Compare commits
12
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bb06c1d634 | ||
|
|
db59678e7e | ||
|
|
c86da2c736 | ||
|
|
bae2a19dcf | ||
|
|
057686f59d | ||
|
|
67da56628b | ||
|
|
2325adffa2 | ||
|
|
20cf836ef1 | ||
|
|
2bf69b6f92 | ||
|
|
13583f5ffb | ||
|
|
008ee2099a | ||
|
|
137f61f2fe |
@@ -28,7 +28,7 @@ def parse_arguments():
|
||||
parser.add_argument(
|
||||
'--image',
|
||||
type=str,
|
||||
default='runpod/pytorch:2.4.0-py3.11-cuda12.4.1-devel-ubuntu22.04',
|
||||
required=True,
|
||||
help='Docker image to use')
|
||||
return parser.parse_args()
|
||||
|
||||
@@ -46,6 +46,16 @@ HEADERS = {
|
||||
|
||||
def create_pod():
|
||||
"""Create a RunPod instance"""
|
||||
# Ensure image name is lowercase (Docker requirement)
|
||||
image_name = args.image.lower()
|
||||
print(f"Using specified image: {image_name}")
|
||||
|
||||
docker_start_cmd = [
|
||||
"bash",
|
||||
"-c",
|
||||
"apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
|
||||
]
|
||||
|
||||
print(f"Creating RunPod instance with GPU: {args.gpu_type}...")
|
||||
payload = {
|
||||
"name": f"fastvideo-{JOB_ID}-{RUN_ID}",
|
||||
@@ -53,8 +63,9 @@ def create_pod():
|
||||
"volumeInGb": args.volume_size,
|
||||
"gpuTypeIds": [args.gpu_type],
|
||||
"gpuCount": args.gpu_count,
|
||||
"imageName": args.image,
|
||||
"allowedCudaVersions": ["12.4"]
|
||||
"imageName": image_name,
|
||||
"allowedCudaVersions": ["12.4"],
|
||||
"dockerStartCmd": docker_start_cmd
|
||||
}
|
||||
|
||||
response = requests.post(PODS_API, headers=HEADERS, json=payload)
|
||||
@@ -91,7 +102,7 @@ def wait_for_pod(pod_id):
|
||||
"Timed out waiting for RunPod to reach RUNNING state")
|
||||
|
||||
# Wait for ports to be assigned
|
||||
max_attempts = 6
|
||||
max_attempts = 50
|
||||
attempts = 0
|
||||
while attempts < max_attempts:
|
||||
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
@@ -108,7 +119,7 @@ def wait_for_pod(pod_id):
|
||||
print(
|
||||
f"Waiting for SSH port and public IP to be available... (attempt {attempts+1}/{max_attempts})"
|
||||
)
|
||||
time.sleep(10)
|
||||
time.sleep(20)
|
||||
attempts += 1
|
||||
|
||||
if attempts >= max_attempts:
|
||||
@@ -145,16 +156,15 @@ def execute_command(pod_id):
|
||||
]
|
||||
subprocess.run(scp_command, check=True)
|
||||
|
||||
# For custom image, we can use the pre-configured environment
|
||||
setup_steps = [
|
||||
"cd /workspace",
|
||||
"wget -q https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh",
|
||||
"bash Miniconda3-latest-Linux-x86_64.sh -b -p $HOME/miniconda3",
|
||||
"source $HOME/miniconda3/bin/activate",
|
||||
"conda create --name venv python=3.10.0 -y", "conda activate venv",
|
||||
"mkdir -p /workspace/repo",
|
||||
"tar -xzf /tmp/repo.tar.gz --no-same-owner -C /workspace/",
|
||||
f"cd /workspace/{repo_name}", args.test_command
|
||||
f"cd /workspace/{repo_name}",
|
||||
"source /opt/conda/etc/profile.d/conda.sh",
|
||||
"conda activate fastvideo-dev",
|
||||
args.test_command
|
||||
]
|
||||
|
||||
remote_command = " && ".join(setup_steps)
|
||||
|
||||
ssh_command = [
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
name: Build and Push Docker Image
|
||||
|
||||
on:
|
||||
workflow_dispatch: # Only manual triggers
|
||||
|
||||
jobs:
|
||||
build-and-push:
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
# Display initial space
|
||||
echo "Initial disk space:"
|
||||
df -h
|
||||
|
||||
# Remove large directories directly
|
||||
sudo rm -rf /usr/share/dotnet
|
||||
sudo rm -rf /usr/local/lib/android
|
||||
sudo rm -rf /opt/ghc
|
||||
sudo rm -rf /usr/local/share/boost
|
||||
sudo rm -rf /usr/share/swift
|
||||
sudo rm -rf /usr/local/lib/node_modules
|
||||
sudo rm -rf /usr/local/share/powershell
|
||||
sudo rm -rf /usr/share/rust
|
||||
sudo rm -rf /usr/local/.ghcup
|
||||
|
||||
# Remove cached files
|
||||
sudo rm -rf /var/lib/apt/lists/*
|
||||
sudo rm -rf /var/cache/apt/archives/*
|
||||
|
||||
# Clean Docker
|
||||
docker system prune -af --volumes
|
||||
|
||||
# Display available space after cleanup
|
||||
echo "Disk space after cleanup:"
|
||||
df -h
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v3
|
||||
|
||||
- name: Login to GitHub Container Registry
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
registry: ghcr.io
|
||||
username: ${{ github.repository_owner }}
|
||||
password: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Extract metadata for Docker
|
||||
id: meta
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: ghcr.io/${{ github.repository }}/fastvideo-dev
|
||||
tags: |
|
||||
type=raw,value=latest
|
||||
type=sha,format=short
|
||||
|
||||
- name: Build and push Docker image
|
||||
id: build-push
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
push: true
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
|
||||
- name: Success message
|
||||
run: |
|
||||
echo "✅ Image successfully built and pushed to ghcr.io/${{ github.repository }}/fastvideo-dev:latest"
|
||||
echo "To run tests with this image, manually trigger the 'Run Tests' workflow."
|
||||
@@ -36,9 +36,13 @@ defaults:
|
||||
shell: bash
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
uses: ./.github/workflows/pre-commit.yml
|
||||
|
||||
# Build job
|
||||
build:
|
||||
runs-on: ubuntu-latest
|
||||
needs: pre-commit
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
|
||||
@@ -6,6 +6,7 @@ on:
|
||||
- main
|
||||
paths:
|
||||
- 'pyproject.toml' # Trigger when pyproject.toml changes
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
check-version-change:
|
||||
@@ -41,7 +42,7 @@ jobs:
|
||||
|
||||
build-publish-main:
|
||||
needs: check-version-change
|
||||
if: needs.check-version-change.outputs.version-changed == 'true'
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
+133
-11
@@ -14,17 +14,35 @@ on:
|
||||
- ".github/workflows/pr-test.yml"
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
custom_image:
|
||||
description: "Custom image from this repository (default: fastvideo-dev:latest)"
|
||||
required: false
|
||||
default: "fastvideo-dev:latest"
|
||||
type: string
|
||||
run_encoder_test:
|
||||
description: "Run encoder-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_vae_test:
|
||||
description: "Run vae-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_transformer_test:
|
||||
description: "Run transformer-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_ssim_test:
|
||||
description: "Run ssim-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
|
||||
env:
|
||||
PYTHONUNBUFFERED: "1"
|
||||
|
||||
concurrency:
|
||||
group: pr-test-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
@@ -39,6 +57,8 @@ jobs:
|
||||
if: ${{ github.event.pull_request.draft == false || github.event_name == 'workflow_dispatch' }}
|
||||
outputs:
|
||||
encoder-test: ${{ steps.filter.outputs.encoder-test }}
|
||||
vae-test: ${{ steps.filter.outputs.vae-test }}
|
||||
transformer-test: ${{ steps.filter.outputs.transformer-test }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
@@ -49,6 +69,14 @@ jobs:
|
||||
- 'fastvideo/v1/models/encoders/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/encoders/**'
|
||||
vae-test:
|
||||
- 'fastvideo/v1/models/vaes/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/vaes/**'
|
||||
transformer-test:
|
||||
- 'fastvideo/v1/models/dits/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/transformers/**'
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
@@ -87,9 +115,8 @@ jobs:
|
||||
--gpu-type "NVIDIA A40"
|
||||
--gpu-count 1
|
||||
--volume-size 100
|
||||
--test-command "pip install -e .[test] &&
|
||||
pip install flash-attn==2.7.0.post2 --no-build-isolation &&
|
||||
pytest ./fastvideo/v1/tests/encoders -s"
|
||||
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
--test-command "pip install -e . && pytest ./fastvideo/v1/tests/encoders -s"
|
||||
|
||||
- name: Terminate RunPod Instances
|
||||
if: ${{ always() }}
|
||||
@@ -99,6 +126,102 @@ jobs:
|
||||
JOB_ID: "encoder-test"
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
vae-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.vae-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_vae_test == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
environment: runpod-runners
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
- name: Set up SSH key
|
||||
run: |
|
||||
mkdir -p ~/.ssh
|
||||
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
|
||||
chmod 600 ~/.ssh/id_rsa
|
||||
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
|
||||
|
||||
- name: Install dependencies
|
||||
run: pip install requests
|
||||
|
||||
- name: Run tests on RunPod
|
||||
env:
|
||||
JOB_ID: "vae-test"
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
timeout-minutes: 30
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
--gpu-type "NVIDIA A40"
|
||||
--gpu-count 1
|
||||
--volume-size 100
|
||||
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
--test-command "pip install -e . && pytest ./fastvideo/v1/tests/vaes -s"
|
||||
|
||||
- name: Terminate RunPod Instances
|
||||
if: ${{ always() }}
|
||||
env:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
JOB_ID: "vae-test"
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
transformer-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.transformer-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_transformer_test == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
environment: runpod-runners
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
- name: Set up SSH key
|
||||
run: |
|
||||
mkdir -p ~/.ssh
|
||||
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
|
||||
chmod 600 ~/.ssh/id_rsa
|
||||
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
|
||||
|
||||
- name: Install dependencies
|
||||
run: pip install requests
|
||||
|
||||
- name: Run tests on RunPod
|
||||
env:
|
||||
JOB_ID: "transformer-test"
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
timeout-minutes: 30
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
--gpu-type "NVIDIA L40S"
|
||||
--gpu-count 1
|
||||
--volume-size 100
|
||||
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
--test-command "pip install -e . && pytest ./fastvideo/v1/tests/transformers -s"
|
||||
|
||||
- name: Terminate RunPod Instances
|
||||
if: ${{ always() }}
|
||||
env:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
JOB_ID: "transformer-test"
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
ssim-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
@@ -130,16 +253,15 @@ jobs:
|
||||
JOB_ID: "ssim-test"
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
timeout-minutes: 30
|
||||
timeout-minutes: 45
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
--gpu-type "NVIDIA A40"
|
||||
--gpu-count 2
|
||||
--disk-size 100
|
||||
--volume-size 100
|
||||
--test-command "pip install -e .[test] &&
|
||||
pip install flash-attn==2.7.0.post2 --no-build-isolation &&
|
||||
pytest ./fastvideo/v1/tests/ssim -vs"
|
||||
--disk-size 200
|
||||
--volume-size 200
|
||||
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
--test-command "pip install -e . && pytest ./fastvideo/v1/tests/ssim -vs"
|
||||
|
||||
- name: Terminate RunPod Instances
|
||||
if: ${{ always() }}
|
||||
@@ -150,7 +272,7 @@ jobs:
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
runpod-cleanup:
|
||||
needs: [encoder-test, ssim-test] # Add other jobs to this list as you create them
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
|
||||
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
@@ -167,7 +289,7 @@ jobs:
|
||||
|
||||
- name: Cleanup all RunPod instances
|
||||
env:
|
||||
JOB_IDS: '["encoder-test", "ssim-test"]' # JSON array of job IDs
|
||||
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test"]' # JSON array of job IDs
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
@@ -6,6 +6,7 @@ on:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/sliding_tile_attention/setup.py"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
check-version-change:
|
||||
@@ -43,7 +44,7 @@ jobs:
|
||||
build_wheels:
|
||||
name: Build Wheel
|
||||
needs: check-version-change
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' }}
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ${{ matrix.os }}
|
||||
|
||||
strategy:
|
||||
@@ -57,6 +58,29 @@ jobs:
|
||||
cuda-version: ['12.4.1', '12.5.1', '12.6.3']
|
||||
|
||||
steps:
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
echo "Initial disk space:"
|
||||
df -h
|
||||
|
||||
# Remove large directories
|
||||
sudo rm -rf /usr/share/dotnet
|
||||
sudo rm -rf /usr/local/lib/android
|
||||
sudo rm -rf /opt/ghc
|
||||
sudo rm -rf /usr/local/share/boost
|
||||
sudo rm -rf /usr/share/swift
|
||||
sudo rm -rf /usr/local/lib/node_modules
|
||||
sudo rm -rf /usr/local/share/powershell
|
||||
sudo rm -rf /usr/share/rust
|
||||
sudo rm -rf /usr/local/.ghcup
|
||||
|
||||
# Remove cached files
|
||||
sudo rm -rf /var/lib/apt/lists/*
|
||||
sudo rm -rf /var/cache/apt/archives/*
|
||||
|
||||
echo "Disk space after cleanup:"
|
||||
df -h
|
||||
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
|
||||
@@ -145,7 +169,7 @@ jobs:
|
||||
publish_package:
|
||||
name: Publish package
|
||||
needs: [build_wheels, check-version-change]
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' }}
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ubuntu-22.04
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
+29
-7
@@ -3,14 +3,11 @@ __pycache__
|
||||
*.pth
|
||||
UCF-101/
|
||||
results/
|
||||
build/
|
||||
fastvideo.egg-info/
|
||||
wandb/
|
||||
*.ipynb
|
||||
*.jpg
|
||||
*.safetensors
|
||||
*.mp4
|
||||
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
|
||||
*.png
|
||||
*.gif
|
||||
*.pth
|
||||
@@ -26,11 +23,36 @@ outputs_video
|
||||
sbatch.sh
|
||||
*.out
|
||||
env
|
||||
dist/
|
||||
*.o
|
||||
**/build/
|
||||
**.egg-info
|
||||
**.pyc
|
||||
**.egg
|
||||
**.txt
|
||||
**.json
|
||||
**.json
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
dist/
|
||||
*.egg-info/
|
||||
*.egg
|
||||
eggs/
|
||||
.eggs/
|
||||
|
||||
# Sphinx documentation
|
||||
docs/_build/
|
||||
docs/source/getting_started/examples/
|
||||
|
||||
# VSCode
|
||||
.vscode/
|
||||
|
||||
# DS Store
|
||||
.DS_Store
|
||||
|
||||
# vim swap files
|
||||
*.swo
|
||||
*.swp
|
||||
|
||||
# Python pickle files
|
||||
*.pkl
|
||||
|
||||
# Reference videos
|
||||
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
|
||||
@@ -40,10 +40,10 @@ repos:
|
||||
- id: codespell
|
||||
additional_dependencies: ['tomli']
|
||||
args: ['--toml', 'pyproject.toml']
|
||||
- repo: https://github.com/PyCQA/isort
|
||||
rev: 0a0b7a830386ba6a31c2ec8316849ae4d1b8240d # 6.0.0
|
||||
hooks:
|
||||
- id: isort
|
||||
# - repo: https://github.com/PyCQA/isort
|
||||
# rev: 0a0b7a830386ba6a31c2ec8316849ae4d1b8240d # 6.0.0
|
||||
# hooks:
|
||||
# - id: isort
|
||||
- repo: https://github.com/jackdewinter/pymarkdown
|
||||
rev: v0.9.29
|
||||
hooks:
|
||||
@@ -66,7 +66,7 @@ repos:
|
||||
entry: bash
|
||||
args:
|
||||
- -c
|
||||
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
|
||||
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
|
||||
language: system
|
||||
always_run: true
|
||||
pass_filenames: false
|
||||
|
||||
@@ -1,21 +0,0 @@
|
||||
# Read the Docs configuration file
|
||||
# See https://docs.readthedocs.io/en/stable/config-file/v2.html for details
|
||||
|
||||
version: 2
|
||||
|
||||
build:
|
||||
os: ubuntu-22.04
|
||||
tools:
|
||||
python: "3.12"
|
||||
|
||||
sphinx:
|
||||
configuration: docs/source/conf.py
|
||||
fail_on_warning: true
|
||||
|
||||
# If using Sphinx, optionally build your docs in additional formats such as PDF
|
||||
formats: []
|
||||
|
||||
# Optionally declare the Python requirements required to build your docs
|
||||
python:
|
||||
install:
|
||||
- requirements: docs/requirements-docs.txt
|
||||
+37
@@ -0,0 +1,37 @@
|
||||
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
wget \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
|
||||
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
|
||||
rm Miniconda3-latest-Linux-x86_64.sh
|
||||
|
||||
ENV PATH=/opt/conda/bin:$PATH
|
||||
|
||||
RUN conda create --name fastvideo-dev python=3.10.0 -y
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject.toml ./
|
||||
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.0.post2 --no-build-isolation && \
|
||||
conda clean -afy
|
||||
|
||||
COPY . .
|
||||
|
||||
EXPOSE 22
|
||||
@@ -5,7 +5,7 @@
|
||||
FastVideo is a lightweight framework for accelerating large video diffusion models.
|
||||
|
||||
<p align="center">
|
||||
🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank">FastHunyuan</a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank">FastMochi</a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> Slack </a>
|
||||
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> <b>Slack</b> </a> |
|
||||
</p>
|
||||
|
||||
https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1
|
||||
@@ -44,6 +44,12 @@ pip install flash-attn==2.7.0.post2
|
||||
|
||||
To try Sliding Tile Attention (optional), please follow the instruction in [csrc/sliding_tile_attention/README.md](csrc/sliding_tile_attention/README.md) to install STA.
|
||||
|
||||
You can also install the Sliding Tile Attention package using
|
||||
|
||||
```
|
||||
pip install st_attn==0.0.3
|
||||
```
|
||||
|
||||
## 🚀 Inference
|
||||
### Inference StepVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
@@ -9,7 +9,7 @@ target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "st_attn"
|
||||
VERSION = "0.0.2"
|
||||
VERSION = "0.0.3"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Sliding Tile Atteniton Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/sliding_tile_attention"
|
||||
|
||||
@@ -446,7 +446,7 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
auto threads = NUM_WORKERS * kittens::WARP_THREADS;
|
||||
if (has_text) {
|
||||
// TORCH_CHECK(seq_len % (CONSUMER_WARPGROUPS*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 192");
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4-2), qo_heads, batch);
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4)-2, qo_heads, batch);
|
||||
dim3 grid_text(2, qo_heads, batch);
|
||||
if (!process_text) {
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
(developer-guide)=
|
||||
|
||||
# Contributing to FastVideo
|
||||
|
||||
Thank you for your interest in contributing to FastVideo. We want to make the process as smooth for you as possible and this is a guide to help get you started!
|
||||
@@ -27,7 +29,7 @@ conda activate fastvideo
|
||||
Clone the FastVideo repository and go to the FastVideo directory:
|
||||
|
||||
```
|
||||
git clone https://github.com/vllm-project/vllm.git && cd vllm
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
|
||||
```
|
||||
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/docs/source/generate_examples.py
|
||||
|
||||
import itertools
|
||||
import re
|
||||
@@ -177,28 +178,28 @@ def generate_examples():
|
||||
# Category indices stored in reverse order because they are inserted into
|
||||
# examples_index.documents at index 0 in order
|
||||
category_indices = {
|
||||
"other":
|
||||
# "other":
|
||||
# Index(
|
||||
# path=EXAMPLE_DOC_DIR / "examples_other_index.md",
|
||||
# title="Other",
|
||||
# description=
|
||||
# "Other examples that don't strongly fit into the online or offline serving categories.", # noqa: E501
|
||||
# caption="Examples",
|
||||
# ),
|
||||
# "online_serving":
|
||||
# Index(
|
||||
# path=EXAMPLE_DOC_DIR / "examples_online_serving_index.md",
|
||||
# title="Online Serving",
|
||||
# description=
|
||||
# "Online serving examples demonstrate how to use FastVideo in an online setting, where the model is queried for predictions in real-time.", # noqa: E501
|
||||
# caption="Examples",
|
||||
# ),
|
||||
"inference":
|
||||
Index(
|
||||
path=EXAMPLE_DOC_DIR / "examples_other_index.md",
|
||||
title="Other",
|
||||
path=EXAMPLE_DOC_DIR / "examples_inference_index.md",
|
||||
title="Inference",
|
||||
description=
|
||||
"Other examples that don't strongly fit into the online or offline serving categories.", # noqa: E501
|
||||
caption="Examples",
|
||||
),
|
||||
"online_serving":
|
||||
Index(
|
||||
path=EXAMPLE_DOC_DIR / "examples_online_serving_index.md",
|
||||
title="Online Serving",
|
||||
description=
|
||||
"Online serving examples demonstrate how to use FastVideo in an online setting, where the model is queried for predictions in real-time.", # noqa: E501
|
||||
caption="Examples",
|
||||
),
|
||||
"offline_inference":
|
||||
Index(
|
||||
path=EXAMPLE_DOC_DIR / "examples_offline_inference_index.md",
|
||||
title="Offline Inference",
|
||||
description=
|
||||
"Offline inference examples demonstrate how to use FastVideo in an offline setting, where the model is queried for predictions in batches. We recommend starting with <project:basic.md>.", # noqa: E501
|
||||
"Inference examples demonstrate how to use FastVideo in an offline setting, where the model is queried for predictions in batches. We recommend starting with <project:basic.md>.", # noqa: E501
|
||||
caption="Examples",
|
||||
),
|
||||
}
|
||||
|
||||
@@ -1,10 +0,0 @@
|
||||
# Examples
|
||||
|
||||
A collection of examples demonstrating usage of FastVideo.
|
||||
All documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.
|
||||
|
||||
:::{toctree}
|
||||
:caption: Examples
|
||||
:maxdepth: 2
|
||||
|
||||
:::
|
||||
@@ -1,10 +1,85 @@
|
||||
(fastvideo-installation)=
|
||||
|
||||
# 🔧 Installation
|
||||
The code is tested on Python 3.10.0, CUDA 12.4 and H100.
|
||||
|
||||
```
|
||||
./env_setup.sh fastvideo
|
||||
FastVideo currently only supports Linux and CUDA GPUs. The code is tested on Python 3.10.0 and CUDA 12.4, primarily with NVIDIA H100 GPUs.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
- CUDA 12.4 installed and supported
|
||||
- Linux operating system
|
||||
|
||||
## Installation Options
|
||||
|
||||
### Option 1: Quick Install
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
To try Sliding Tile Attention (optional), please follow the instruction in [here](#sta-installation) to install STA.
|
||||
### Option 2: Installation from Source
|
||||
|
||||
#### 1. Install Miniconda (if not already installed)
|
||||
|
||||
```bash
|
||||
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
|
||||
bash Miniconda3-latest-Linux-x86_64.sh
|
||||
source ~/.bashrc
|
||||
```
|
||||
|
||||
#### 2. Create and activate a Conda environment for FastVideo
|
||||
|
||||
```bash
|
||||
conda create -n fastvideo python=3.10 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
#### 3. Clone the FastVideo repository
|
||||
|
||||
```bash
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
```
|
||||
|
||||
#### 4. Install FastVideo
|
||||
|
||||
Basic installation:
|
||||
|
||||
```bash
|
||||
pip install -e .
|
||||
```
|
||||
|
||||
## Optional Dependencies
|
||||
|
||||
### Flash Attention
|
||||
|
||||
```bash
|
||||
pip install flash-attn==2.7.0.post2 --no-build-isolation
|
||||
```
|
||||
|
||||
### Sliding Tile Attention (STA)
|
||||
|
||||
To try Sliding Tile Attention (optional), please follow the instructions in [csrc/sliding_tile_attention/README.md](#sta-installation) to install STA.
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
If you're planning to contribute to FastVideo please see the following page:
|
||||
[Contributor Guide](#developer-guide)
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
### For Basic Inference
|
||||
- NVIDIA GPU with CUDA support
|
||||
- Minimum 20GB VRAM for quantized models (e.g., single RTX 4090)
|
||||
|
||||
### For Lora Finetuning
|
||||
- 40GB GPU memory each for 2 GPUs with lora
|
||||
- 30GB GPU memory each for 2 GPUs with CPU offload and lora
|
||||
|
||||
### For Full Finetuning/Distillation
|
||||
- Multiple high-memory GPUs recommended (e.g., H100)
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
If you encounter any issues during installation, please open an issue on our [GitHub repository](https://github.com/hao-ai-lab/FastVideo).
|
||||
|
||||
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg) for additional support.
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
(wanvideo)=
|
||||
|
||||
# WanVideo
|
||||
## Inference T2V with WanVideo
|
||||
First, download the model:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-T2V-1.3B-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
|
||||
```
|
||||
|
||||
or
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-T2V-14B-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
|
||||
```
|
||||
|
||||
Then run the inference using:
|
||||
|
||||
```bash
|
||||
sh scripts/inference/v1_inference_wan.sh
|
||||
```
|
||||
|
||||
Remember to set `MODEL_BASE` and `num_gpus` accordingly.
|
||||
|
||||
## Inference I2V with WanVideo
|
||||
First, download the model:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-I2V-14B-480P-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
|
||||
```
|
||||
|
||||
or
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-I2V-14B-720P-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
|
||||
```
|
||||
|
||||
Then run the inference using:
|
||||
|
||||
```bash
|
||||
sh scripts/inference/v1_inference_wan_i2v.sh
|
||||
```
|
||||
|
||||
Remember to set `MODEL_BASE` and `num_gpus` accordingly.
|
||||
@@ -0,0 +1,3 @@
|
||||
# Basic
|
||||
|
||||
The class provides the main python interface for using FastVideo's inference pipeline.
|
||||
@@ -0,0 +1 @@
|
||||
print('Hello, world!')
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -7,7 +7,7 @@ from typing import (TYPE_CHECKING, Any, Dict, Generic, Optional, Protocol, Set,
|
||||
Type, TypeVar)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
import torch
|
||||
@@ -154,7 +154,7 @@ class AttentionMetadataBuilder(ABC, Generic[T]):
|
||||
self,
|
||||
current_timestep: int,
|
||||
forward_batch: "ForwardBatch",
|
||||
inference_args: "InferenceArgs",
|
||||
fastvideo_args: "FastVideoArgs",
|
||||
) -> T:
|
||||
"""Build attention metadata with on-device tensors."""
|
||||
raise NotImplementedError
|
||||
@@ -186,9 +186,10 @@ class AttentionImpl(ABC, Generic[T]):
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
softmax_scale: float,
|
||||
dropout_rate: float = 0.0,
|
||||
causal: bool = False,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@@ -3,7 +3,16 @@
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
from flash_attn import flash_attn_func
|
||||
from flash_attn import flash_attn_func as flash_attn_2_func
|
||||
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_3_func
|
||||
|
||||
# flash_attn 3 has slightly different API: it returns lse by default
|
||||
flash_attn_func = lambda q, k, v, softmax_scale, causal: flash_attn_3_func(
|
||||
q, k, v, softmax_scale, causal)[0]
|
||||
except ImportError:
|
||||
flash_attn_func = flash_attn_2_func
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
@@ -45,12 +54,12 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
dropout_rate: float,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
self.dropout_rate = dropout_rate
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
|
||||
@@ -61,10 +70,10 @@ class FlashAttentionImpl(AttentionImpl):
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
):
|
||||
output = flash_attn_func(query,
|
||||
key,
|
||||
value,
|
||||
dropout_p=self.dropout_rate,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal)
|
||||
output = flash_attn_func(
|
||||
query, # type: ignore[no-untyped-call]
|
||||
key,
|
||||
value,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal)
|
||||
return output
|
||||
|
||||
@@ -38,14 +38,15 @@ class SDPAImpl(AttentionImpl):
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
dropout_rate: float,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
self.dropout_rate = dropout_rate
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
self.dropout = extra_impl_args.get("dropout_p", 0.0)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -60,7 +61,7 @@ class SDPAImpl(AttentionImpl):
|
||||
value = value.transpose(1, 2)
|
||||
attn_kwargs = {
|
||||
"attn_mask": None,
|
||||
"dropout_p": self.dropout_rate,
|
||||
"dropout_p": self.dropout,
|
||||
"is_causal": self.causal,
|
||||
"scale": self.softmax_scale
|
||||
}
|
||||
|
||||
@@ -12,7 +12,7 @@ from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.distributed import get_sp_group
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
@@ -62,7 +62,7 @@ class SlidingTileAttentionBackend(AttentionBackend):
|
||||
|
||||
@dataclass
|
||||
class SlidingTileAttentionMetadata(AttentionMetadata):
|
||||
text_length: int
|
||||
current_timestep: int
|
||||
|
||||
|
||||
class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
@@ -77,13 +77,10 @@ class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
self,
|
||||
current_timestep: int,
|
||||
forward_batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> SlidingTileAttentionMetadata:
|
||||
|
||||
return SlidingTileAttentionMetadata(
|
||||
current_timestep=current_timestep,
|
||||
text_length=forward_batch.attention_mask.sum(),
|
||||
)
|
||||
return SlidingTileAttentionMetadata(current_timestep=current_timestep, )
|
||||
|
||||
|
||||
class SlidingTileAttentionImpl(AttentionImpl):
|
||||
@@ -92,10 +89,11 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
dropout_rate: float,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args,
|
||||
) -> None:
|
||||
# TODO(will-refactor): for now this is the mask strategy, but maybe we should
|
||||
# have a more general config for STA?
|
||||
@@ -107,7 +105,7 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
mask_strategy = json.load(f)
|
||||
|
||||
mask_strategy = dict_to_3d_list(mask_strategy)
|
||||
|
||||
self.prefix = prefix
|
||||
self.mask_strategy = mask_strategy
|
||||
sp_group = get_sp_group()
|
||||
self.sp_size = sp_group.world_size
|
||||
@@ -172,8 +170,11 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
assert self.mask_strategy[
|
||||
0] is not None, "mask_strategy[0] cannot be None for SlidingTileAttention"
|
||||
|
||||
text_length = attn_metadata.text_length
|
||||
|
||||
timestep = attn_metadata.current_timestep
|
||||
# pattern:'.double_blocks.0.attn.impl' or '.single_blocks.0.attn.impl'
|
||||
layer_idx = int(self.prefix.split('.')[-3])
|
||||
# TODO: remove hardcode
|
||||
text_length = q.shape[1] - (30 * 48 * 80)
|
||||
query = q.transpose(1, 2)
|
||||
key = k.transpose(1, 2)
|
||||
value = v.transpose(1, 2)
|
||||
@@ -183,13 +184,10 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
current_rank = sp_group.rank_in_group
|
||||
start_head = current_rank * head_num
|
||||
windows = [
|
||||
self.mask_strategy[head_idx + start_head]
|
||||
self.mask_strategy[timestep][layer_idx][head_idx + start_head]
|
||||
for head_idx in range(head_num)
|
||||
]
|
||||
|
||||
hidden_states = sliding_tile_attention(query, key, value, windows,
|
||||
text_length).transpose(1, 2)
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from typing import Optional
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
@@ -12,6 +12,7 @@ from fastvideo.v1.distributed.communication_op import (
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
class DistributedAttention(nn.Module):
|
||||
@@ -22,13 +23,12 @@ class DistributedAttention(nn.Module):
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
dropout_rate: float = 0.0,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: Optional[List[_Backend]] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
# self.dropout_rate = dropout_rate
|
||||
# self.causal = causal
|
||||
if softmax_scale is None:
|
||||
self.softmax_scale = head_size**-0.5
|
||||
else:
|
||||
@@ -38,14 +38,17 @@ class DistributedAttention(nn.Module):
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = torch.get_default_dtype()
|
||||
attn_backend = get_attn_backend(head_size, dtype, distributed=True)
|
||||
attn_backend = get_attn_backend(
|
||||
head_size,
|
||||
dtype,
|
||||
supported_attention_backends=supported_attention_backends)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
dropout_rate=dropout_rate,
|
||||
causal=causal,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
prefix=f"{prefix}.impl",
|
||||
**extra_impl_args)
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
@@ -97,7 +100,6 @@ class DistributedAttention(nn.Module):
|
||||
qkv = sequence_model_parallel_all_to_all_4D(qkv,
|
||||
scatter_dim=2,
|
||||
gather_dim=1)
|
||||
|
||||
# Apply backend-specific preprocess_qkv
|
||||
qkv = self.impl.preprocess_qkv(qkv, ctx_attn_metadata)
|
||||
|
||||
@@ -124,8 +126,7 @@ class DistributedAttention(nn.Module):
|
||||
output = output[:, :seq_len * world_size]
|
||||
# TODO: make this asynchronous
|
||||
replicated_output = sequence_model_parallel_all_gather(
|
||||
replicated_output, dim=2)
|
||||
|
||||
replicated_output.contiguous(), dim=2)
|
||||
# Apply backend-specific postprocess_output
|
||||
output = self.impl.postprocess_output(output, ctx_attn_metadata)
|
||||
|
||||
@@ -143,13 +144,11 @@ class LocalAttention(nn.Module):
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
dropout_rate: float = 0.0,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: Optional[List[_Backend]] = None,
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
# self.dropout_rate = dropout_rate
|
||||
# self.causal = causal
|
||||
if softmax_scale is None:
|
||||
self.softmax_scale = head_size**-0.5
|
||||
else:
|
||||
@@ -158,11 +157,13 @@ class LocalAttention(nn.Module):
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = torch.get_default_dtype()
|
||||
attn_backend = get_attn_backend(head_size, dtype, distributed=False)
|
||||
attn_backend = get_attn_backend(
|
||||
head_size,
|
||||
dtype,
|
||||
supported_attention_backends=supported_attention_backends)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
dropout_rate=dropout_rate,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
causal=causal,
|
||||
|
||||
@@ -3,8 +3,7 @@
|
||||
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from functools import cache
|
||||
from typing import Generator, Optional, Type, cast
|
||||
from typing import Generator, List, Optional, Type, cast
|
||||
|
||||
import torch
|
||||
|
||||
@@ -82,29 +81,15 @@ def get_global_forced_attn_backend() -> Optional[_Backend]:
|
||||
def get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
distributed: bool,
|
||||
) -> Type[AttentionBackend]:
|
||||
"""Selects which attention backend to use and lazily imports it."""
|
||||
# Accessing envs.* behind an @lru_cache decorator can cause the wrong
|
||||
# value to be returned from the cache if the value changes between calls.
|
||||
return _cached_get_attn_backend(
|
||||
head_size=head_size,
|
||||
dtype=dtype,
|
||||
distributed=distributed,
|
||||
)
|
||||
|
||||
|
||||
@cache
|
||||
def _cached_get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
distributed: bool,
|
||||
supported_attention_backends: Optional[List[_Backend]] = None,
|
||||
) -> Type[AttentionBackend]:
|
||||
# Check whether a particular choice of backend was
|
||||
# previously forced.
|
||||
#
|
||||
# THIS SELECTION OVERRIDES THE FASTVIDEO_ATTENTION_BACKEND
|
||||
# ENVIRONMENT VARIABLE.
|
||||
if not supported_attention_backends:
|
||||
raise ValueError("supported_attention_backends is empty")
|
||||
selected_backend = None
|
||||
backend_by_global_setting: Optional[_Backend] = (
|
||||
get_global_forced_attn_backend())
|
||||
@@ -117,8 +102,10 @@ def _cached_get_attn_backend(
|
||||
selected_backend = backend_name_to_enum(backend_by_env_var)
|
||||
|
||||
# get device-specific attn_backend
|
||||
if selected_backend not in supported_attention_backends:
|
||||
selected_backend = None
|
||||
attention_cls = current_platform.get_attn_backend_cls(
|
||||
selected_backend, head_size, dtype, distributed)
|
||||
selected_backend, head_size, dtype)
|
||||
if not attention_cls:
|
||||
raise ValueError(
|
||||
f"Invalid attention backend for {current_platform.device_name}")
|
||||
|
||||
@@ -0,0 +1,8 @@
|
||||
from fastvideo.v1.configs.hunyuan import HunyuanConfig, FastHunyuanConfig
|
||||
from fastvideo.v1.configs.wan import WanT2V480PConfig, WanI2V480PConfig
|
||||
from fastvideo.v1.configs.base import BaseConfig, SlidingTileAttnConfig
|
||||
|
||||
__all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "BaseConfig", "SlidingTileAttnConfig",
|
||||
"WanT2V480PConfig", "WanI2V480PConfig"
|
||||
]
|
||||
@@ -0,0 +1,70 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Dict, Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class BaseConfig:
|
||||
"""Base configuration for all pipeline architectures."""
|
||||
|
||||
# Video parameters
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
num_frames: int = 125
|
||||
fps: int = 24
|
||||
|
||||
# Video generation parameters
|
||||
num_inference_steps: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
seed: int = 1024
|
||||
guidance_rescale: float = 0.0
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: Optional[float] = None
|
||||
use_cpu_offload: bool = False
|
||||
disable_autocast: bool = False
|
||||
|
||||
# Model configuration
|
||||
precision: str = "bf16"
|
||||
|
||||
# VAE configuration
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = True
|
||||
vae_scale_factor: Optional[int] = None
|
||||
|
||||
# DiT configuration
|
||||
num_channels_latents: Optional[int] = None
|
||||
|
||||
# Image encoder configuration
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precision: str = "fp16"
|
||||
text_len: int = -1
|
||||
hidden_state_skip_layer: int = 0
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
neg_prompt: Optional[str] = None
|
||||
|
||||
# Additional parameters can be added as a dict
|
||||
extra_params: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class SlidingTileAttnConfig(BaseConfig):
|
||||
"""Configuration for sliding tile attention."""
|
||||
|
||||
# Override any BaseConfig defaults as needed
|
||||
# Add sliding tile specific parameters
|
||||
window_size: int = 16
|
||||
stride: int = 8
|
||||
|
||||
# You can provide custom defaults for inherited fields
|
||||
height: int = 576
|
||||
width: int = 1024
|
||||
|
||||
# Additional configuration specific to sliding tile attention
|
||||
pad_to_square: bool = False
|
||||
use_overlap_optimization: bool = True
|
||||
@@ -0,0 +1,39 @@
|
||||
from dataclasses import dataclass
|
||||
from fastvideo.v1.configs.base import BaseConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class HunyuanConfig(BaseConfig):
|
||||
"""Base configuration for HunYuan pipeline architecture."""
|
||||
|
||||
# HunyuanConfig-specific parameters with defaults
|
||||
# Denoising stage
|
||||
embedded_cfg_scale: int = 6
|
||||
flow_shift: int = 7
|
||||
num_inference_steps: int = 50
|
||||
|
||||
# Text encoding stage
|
||||
hidden_state_skip_layer: int = 2
|
||||
text_len: int = 256
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precision: str = "fp16"
|
||||
|
||||
# HunyuanConfig-specific added parameters
|
||||
# Secondary text encoder
|
||||
text_encoder_precision_2: str = "fp16"
|
||||
text_len_2: int = 77
|
||||
|
||||
|
||||
@dataclass
|
||||
class FastHunyuanConfig(HunyuanConfig):
|
||||
"""Configuration specifically optimized for FastHunyuan weights."""
|
||||
|
||||
# Override HunyuanConfig defaults
|
||||
num_inference_steps: int = 6
|
||||
flow_shift: int = 17
|
||||
|
||||
# No need to re-specify guidance_scale or embedded_cfg_scale as they
|
||||
# already have the desired values from HunyuanConfig
|
||||
@@ -0,0 +1,77 @@
|
||||
"""Registry for pipeline weight-specific configurations."""
|
||||
|
||||
import os
|
||||
from typing import Dict, Type, Optional, Callable
|
||||
|
||||
from fastvideo.v1.configs.base import BaseConfig
|
||||
from fastvideo.v1.configs.hunyuan import HunyuanConfig, FastHunyuanConfig
|
||||
from fastvideo.v1.configs.wan import WanT2V480PConfig, WanI2V480PConfig
|
||||
|
||||
from fastvideo.v1.utils import maybe_download_model_index, verify_model_config_and_directory
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Registry maps specific model weights to their config classes
|
||||
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[BaseConfig]] = {
|
||||
"FastVideo/FastHunyuan-Diffusers": FastHunyuanConfig,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig
|
||||
# Add other specific weight variants
|
||||
}
|
||||
|
||||
# For determining pipeline type from model ID
|
||||
PIPELINE_DETECTOR: Dict[str, Callable[[str], bool]] = {
|
||||
"hunyuan": lambda id: "hunyuan" in id.lower(),
|
||||
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
|
||||
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
|
||||
# Add other pipeline architecture detectors
|
||||
}
|
||||
|
||||
# Fallback configs when exact match isn't found but architecture is detected
|
||||
PIPELINE_FALLBACK_CONFIG: Dict[str, Type[BaseConfig]] = {
|
||||
"hunyuan":
|
||||
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
|
||||
"wanpipeline":
|
||||
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
|
||||
"wanimagetovideo": WanI2V480PConfig,
|
||||
# Other fallbacks by architecture
|
||||
}
|
||||
|
||||
|
||||
def get_pipeline_config_for_name(
|
||||
pipeline_name_or_path: str) -> Optional[Type[BaseConfig]]:
|
||||
"""Get the appropriate config class for specific pretrained weights."""
|
||||
|
||||
if os.path.exists(pipeline_name_or_path):
|
||||
config = verify_model_config_and_directory(pipeline_name_or_path)
|
||||
logger.warning(
|
||||
"FastVideo may not correctly identify the optimal config for this model, as the local directory may have been renamed."
|
||||
)
|
||||
else:
|
||||
config = maybe_download_model_index(pipeline_name_or_path)
|
||||
|
||||
pipeline_name = config["_class_name"]
|
||||
|
||||
# First try exact match for specific weights
|
||||
if pipeline_name_or_path in WEIGHT_CONFIG_REGISTRY:
|
||||
return WEIGHT_CONFIG_REGISTRY[pipeline_name_or_path]
|
||||
|
||||
# Try partial matches (for local paths that might include the weight ID)
|
||||
for registered_id, config_class in WEIGHT_CONFIG_REGISTRY.items():
|
||||
if registered_id in pipeline_name_or_path:
|
||||
return config_class
|
||||
|
||||
# If no match, try to use the fallback config
|
||||
fallback_config = None
|
||||
print(pipeline_name)
|
||||
# Try to determine pipeline architecture for fallback
|
||||
for pipeline_type, detector in PIPELINE_DETECTOR.items():
|
||||
if detector(pipeline_name.lower()):
|
||||
fallback_config = PIPELINE_FALLBACK_CONFIG.get(pipeline_type)
|
||||
break
|
||||
|
||||
logger.warning("No match found for pipeline %s, using fallback config %s.",
|
||||
pipeline_name_or_path, fallback_config)
|
||||
return fallback_config
|
||||
@@ -0,0 +1,44 @@
|
||||
from dataclasses import dataclass
|
||||
from fastvideo.v1.configs.base import BaseConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanT2V480PConfig(BaseConfig):
|
||||
"""Base configuration for Wan T2V 1.3B pipeline architecture."""
|
||||
|
||||
# WanConfig-specific parameters with defaults
|
||||
# Video parameters
|
||||
height: int = 480
|
||||
width: int = 832
|
||||
num_frames: int = 81
|
||||
fps: int = 16
|
||||
use_cpu_offload: bool = True
|
||||
|
||||
# Denoising stage
|
||||
guidance_scale: float = 3.0
|
||||
neg_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
flow_shift: int = 3
|
||||
num_inference_steps: int = 50
|
||||
|
||||
# Text encoding stage
|
||||
text_len: int = 512
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precision: str = "fp32"
|
||||
|
||||
# WanConfig-specific added parameters
|
||||
|
||||
|
||||
@dataclass
|
||||
class WanI2V480PConfig(WanT2V480PConfig):
|
||||
"""Base configuration for Wan I2V 14B 480P pipeline architecture."""
|
||||
|
||||
# WanConfig-specific parameters with defaults
|
||||
# Denoising stage
|
||||
guidance_scale: float = 5.0
|
||||
num_inference_steps: int = 40
|
||||
|
||||
# Precision for each component
|
||||
image_encoder_precision: str = "fp32"
|
||||
@@ -190,5 +190,5 @@ class DeviceCommunicatorBase:
|
||||
torch.distributed.recv(tensor, self.ranks[src], self.device_group)
|
||||
return tensor
|
||||
|
||||
def destroy(self):
|
||||
def destroy(self) -> None:
|
||||
pass
|
||||
|
||||
@@ -845,12 +845,12 @@ def initialize_model_parallel(
|
||||
group_name="sp")
|
||||
|
||||
|
||||
def get_sequence_model_parallel_world_size():
|
||||
def get_sequence_model_parallel_world_size() -> int:
|
||||
"""Return world size for the sequence model parallel group."""
|
||||
return get_sp_group().world_size
|
||||
|
||||
|
||||
def get_sequence_model_parallel_rank():
|
||||
def get_sequence_model_parallel_rank() -> int:
|
||||
"""Return my rank for the sequence model parallel group."""
|
||||
return get_sp_group().rank_in_group
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
class CLISubcommand:
|
||||
"""Base class for CLI subcommands"""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self) -> None:
|
||||
self.name = ""
|
||||
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
|
||||
@@ -2,11 +2,11 @@
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/serve.py
|
||||
|
||||
import argparse
|
||||
from typing import List
|
||||
from typing import List, cast
|
||||
|
||||
from fastvideo.v1.entrypoints.cli import utils
|
||||
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
@@ -82,9 +82,9 @@ class GenerateSubcommand(CLISubcommand):
|
||||
default=None,
|
||||
help="Port for the master process")
|
||||
|
||||
generate_parser = InferenceArgs.add_cli_args(generate_parser)
|
||||
generate_parser = FastVideoArgs.add_cli_args(generate_parser)
|
||||
|
||||
return generate_parser
|
||||
return cast(FlexibleArgumentParser, generate_parser)
|
||||
|
||||
|
||||
def cmd_init() -> List[CLISubcommand]:
|
||||
|
||||
@@ -3,13 +3,16 @@
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from typing import List, Optional
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def launch_distributed(num_gpus=None, args=None, master_port=None):
|
||||
def launch_distributed(num_gpus: int,
|
||||
args: List[str],
|
||||
master_port: Optional[int] = None) -> int:
|
||||
"""
|
||||
Launch a distributed job with the given arguments
|
||||
|
||||
|
||||
@@ -8,8 +8,6 @@ if TYPE_CHECKING:
|
||||
FASTVIDEO_RINGBUFFER_WARNING_INTERVAL: int = 60
|
||||
FASTVIDEO_NCCL_SO_PATH: Optional[str] = None
|
||||
LD_LIBRARY_PATH: Optional[str] = None
|
||||
FASTVIDEO_USE_TRITON_FLASH_ATTN: bool = False
|
||||
FASTVIDEO_FLASH_ATTN_VERSION: Optional[int] = None
|
||||
LOCAL_RANK: int = 0
|
||||
CUDA_VISIBLE_DEVICES: Optional[str] = None
|
||||
FASTVIDEO_CACHE_ROOT: str = os.path.expanduser("~/.cache/fastvideo")
|
||||
@@ -127,18 +125,6 @@ environment_variables: Dict[str, Callable[[], Any]] = {
|
||||
"LD_LIBRARY_PATH":
|
||||
lambda: os.environ.get("LD_LIBRARY_PATH", None),
|
||||
|
||||
# flag to control if fastvideo should use triton flash attention
|
||||
"FASTVIDEO_USE_TRITON_FLASH_ATTN":
|
||||
lambda:
|
||||
(os.environ.get("FASTVIDEO_USE_TRITON_FLASH_ATTN", "True").lower() in
|
||||
("true", "1")),
|
||||
|
||||
# Force fastvideo to use a specific flash-attention version (2 or 3), only valid
|
||||
# when using the flash-attention backend.
|
||||
"FASTVIDEO_FLASH_ATTN_VERSION":
|
||||
lambda: maybe_convert_int(
|
||||
os.environ.get("FASTVIDEO_FLASH_ATTN_VERSION", None)),
|
||||
|
||||
# Internal flag to enable Dynamo fullgraph capture
|
||||
"FASTVIDEO_TEST_DYNAMO_FULLGRAPH_CAPTURE":
|
||||
lambda: bool(
|
||||
|
||||
@@ -4,23 +4,33 @@
|
||||
|
||||
import argparse
|
||||
import dataclasses
|
||||
from contextlib import contextmanager
|
||||
from typing import List, Optional
|
||||
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class InferenceArgs:
|
||||
class FastVideoArgs:
|
||||
# Model and path configuration
|
||||
model_path: str
|
||||
|
||||
# Distributed executor backend
|
||||
distributed_executor_backend: str = "torch"
|
||||
|
||||
inference_mode: bool = True # if False == training mode
|
||||
|
||||
# HuggingFace specific parameters
|
||||
trust_remote_code: bool = False
|
||||
revision: Optional[str] = None
|
||||
|
||||
# Parallelism
|
||||
tp_size: int = 1
|
||||
sp_size: int = 1
|
||||
num_gpus: int = 1
|
||||
tp_size: Optional[int] = None
|
||||
sp_size: Optional[int] = None
|
||||
dist_timeout: Optional[int] = None # timeout for torch.distributed
|
||||
|
||||
# Video generation parameters
|
||||
@@ -31,7 +41,7 @@ class InferenceArgs:
|
||||
guidance_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: int = 7
|
||||
flow_shift: Optional[float] = None
|
||||
|
||||
output_type: str = "pil"
|
||||
|
||||
@@ -42,6 +52,13 @@ class InferenceArgs:
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = False
|
||||
vae_scale_factor: Optional[int] = None
|
||||
|
||||
# DiT configuration
|
||||
num_channels_latents: Optional[int] = None
|
||||
|
||||
# Image encoder configuration
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precision: str = "fp16"
|
||||
@@ -54,14 +71,14 @@ class InferenceArgs:
|
||||
|
||||
# Flow Matching parameters
|
||||
flow_solver: str = "euler"
|
||||
denoise_type: str = "flow"
|
||||
denoise_type: str = "flow" # Deprecated. Will use scheduler_config.json
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
# Scheduler options
|
||||
scheduler_type: str = "euler"
|
||||
scheduler_type: str = "euler" # Deprecated. Will use the param in scheduler_config.json
|
||||
|
||||
neg_prompt: Optional[str] = None
|
||||
num_videos: int = 1
|
||||
@@ -73,6 +90,7 @@ class InferenceArgs:
|
||||
log_level: str = "info"
|
||||
|
||||
# Inference parameters
|
||||
image_path: Optional[str] = None
|
||||
prompt: Optional[str] = None
|
||||
prompt_path: Optional[str] = None
|
||||
output_path: str = "outputs/"
|
||||
@@ -104,40 +122,55 @@ class InferenceArgs:
|
||||
help="Directory containing StepVideo model",
|
||||
)
|
||||
|
||||
# distributed_executor_backend
|
||||
parser.add_argument(
|
||||
"--distributed-executor-backend",
|
||||
type=str,
|
||||
choices=["mp", "ray", "torch"],
|
||||
default=FastVideoArgs.distributed_executor_backend,
|
||||
help="The distributed executor backend to use",
|
||||
)
|
||||
|
||||
# HuggingFace specific parameters
|
||||
parser.add_argument(
|
||||
"--trust-remote-code",
|
||||
action="store_true",
|
||||
default=InferenceArgs.trust_remote_code,
|
||||
default=FastVideoArgs.trust_remote_code,
|
||||
help="Trust remote code when loading HuggingFace models",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--revision",
|
||||
type=str,
|
||||
default=InferenceArgs.revision,
|
||||
default=FastVideoArgs.revision,
|
||||
help=
|
||||
"The specific model version to use (can be a branch name, tag name, or commit id)",
|
||||
)
|
||||
|
||||
# Parallelism
|
||||
parser.add_argument(
|
||||
"--num-gpus",
|
||||
type=int,
|
||||
default=FastVideoArgs.num_gpus,
|
||||
help="The number of GPUs to use.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tensor-parallel-size",
|
||||
"--tp-size",
|
||||
type=int,
|
||||
default=InferenceArgs.tp_size,
|
||||
default=FastVideoArgs.tp_size,
|
||||
help="The tensor parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sequence-parallel-size",
|
||||
"--sp-size",
|
||||
type=int,
|
||||
default=InferenceArgs.sp_size,
|
||||
default=FastVideoArgs.sp_size,
|
||||
help="The sequence parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dist-timeout",
|
||||
type=int,
|
||||
default=InferenceArgs.dist_timeout,
|
||||
default=FastVideoArgs.dist_timeout,
|
||||
help="Set timeout for torch.distributed initialization.",
|
||||
)
|
||||
|
||||
@@ -145,56 +178,56 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--height",
|
||||
type=int,
|
||||
default=InferenceArgs.height,
|
||||
default=FastVideoArgs.height,
|
||||
help="Height of generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--width",
|
||||
type=int,
|
||||
default=InferenceArgs.width,
|
||||
default=FastVideoArgs.width,
|
||||
help="Width of generated video",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-frames",
|
||||
type=int,
|
||||
default=InferenceArgs.num_frames,
|
||||
default=FastVideoArgs.num_frames,
|
||||
help="Number of frames to generate",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-inference-steps",
|
||||
type=int,
|
||||
default=InferenceArgs.num_inference_steps,
|
||||
default=FastVideoArgs.num_inference_steps,
|
||||
help="Number of inference steps",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-scale",
|
||||
type=float,
|
||||
default=InferenceArgs.guidance_scale,
|
||||
default=FastVideoArgs.guidance_scale,
|
||||
help="Guidance scale for classifier-free guidance",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--guidance-rescale",
|
||||
type=float,
|
||||
default=InferenceArgs.guidance_rescale,
|
||||
default=FastVideoArgs.guidance_rescale,
|
||||
help="Guidance rescale for classifier-free guidance",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--embedded-cfg-scale",
|
||||
type=float,
|
||||
default=InferenceArgs.embedded_cfg_scale,
|
||||
default=FastVideoArgs.embedded_cfg_scale,
|
||||
help="Embedded CFG scale",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flow-shift",
|
||||
"--shift",
|
||||
type=int,
|
||||
default=InferenceArgs.flow_shift,
|
||||
type=float,
|
||||
default=FastVideoArgs.flow_shift,
|
||||
help="Flow shift parameter",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output-type",
|
||||
type=str,
|
||||
default=InferenceArgs.output_type,
|
||||
default=FastVideoArgs.output_type,
|
||||
choices=["pil"],
|
||||
help="Output type for the generated video",
|
||||
)
|
||||
@@ -202,7 +235,7 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--precision",
|
||||
type=str,
|
||||
default=InferenceArgs.precision,
|
||||
default=FastVideoArgs.precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for the model",
|
||||
)
|
||||
@@ -211,14 +244,14 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--vae-precision",
|
||||
type=str,
|
||||
default=InferenceArgs.vae_precision,
|
||||
default=FastVideoArgs.vae_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for VAE",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-tiling",
|
||||
action="store_true",
|
||||
default=InferenceArgs.vae_tiling,
|
||||
default=FastVideoArgs.vae_tiling,
|
||||
help="Enable VAE tiling",
|
||||
)
|
||||
parser.add_argument(
|
||||
@@ -230,29 +263,39 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision",
|
||||
type=str,
|
||||
default=InferenceArgs.text_encoder_precision,
|
||||
default=FastVideoArgs.text_encoder_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for text encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-len",
|
||||
type=int,
|
||||
default=InferenceArgs.text_len,
|
||||
default=FastVideoArgs.text_len,
|
||||
help="Maximum text length",
|
||||
)
|
||||
|
||||
# Image encoder config
|
||||
parser.add_argument(
|
||||
"--image-encoder-precision",
|
||||
type=str,
|
||||
default=FastVideoArgs.image_encoder_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for image encoder",
|
||||
)
|
||||
|
||||
# Secondary text encoder
|
||||
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision-2",
|
||||
type=str,
|
||||
default=InferenceArgs.text_encoder_precision_2,
|
||||
default=FastVideoArgs.text_encoder_precision_2,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for secondary text encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--text-len-2",
|
||||
type=int,
|
||||
default=InferenceArgs.text_len_2,
|
||||
default=FastVideoArgs.text_len_2,
|
||||
help="Maximum secondary text length",
|
||||
)
|
||||
|
||||
@@ -260,13 +303,13 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--flow-solver",
|
||||
type=str,
|
||||
default=InferenceArgs.flow_solver,
|
||||
default=FastVideoArgs.flow_solver,
|
||||
help="Solver for flow matching",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--denoise-type",
|
||||
type=str,
|
||||
default=InferenceArgs.denoise_type,
|
||||
default=FastVideoArgs.denoise_type,
|
||||
help="Denoise type for noised inputs",
|
||||
)
|
||||
|
||||
@@ -287,7 +330,7 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--scheduler-type",
|
||||
type=str,
|
||||
default=InferenceArgs.scheduler_type,
|
||||
default=FastVideoArgs.scheduler_type,
|
||||
help="Type of scheduler to use",
|
||||
)
|
||||
|
||||
@@ -295,19 +338,19 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--neg-prompt",
|
||||
type=str,
|
||||
default=InferenceArgs.neg_prompt,
|
||||
default=FastVideoArgs.neg_prompt,
|
||||
help="Negative prompt for sampling",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--num-videos",
|
||||
type=int,
|
||||
default=InferenceArgs.num_videos,
|
||||
default=FastVideoArgs.num_videos,
|
||||
help="Number of videos to generate per prompt",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--fps",
|
||||
type=int,
|
||||
default=InferenceArgs.fps,
|
||||
default=FastVideoArgs.fps,
|
||||
help="Frames per second for output video",
|
||||
)
|
||||
parser.add_argument(
|
||||
@@ -326,7 +369,7 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--log-level",
|
||||
type=str,
|
||||
default=InferenceArgs.log_level,
|
||||
default=FastVideoArgs.log_level,
|
||||
help="The logging level of all loggers.",
|
||||
)
|
||||
|
||||
@@ -343,23 +386,27 @@ class InferenceArgs:
|
||||
help="Path to a text file containing the prompt",
|
||||
)
|
||||
|
||||
parser.add_argument("--image-path",
|
||||
type=str,
|
||||
help="Path to the image for I2V generation")
|
||||
|
||||
parser.add_argument(
|
||||
"--output-path",
|
||||
type=str,
|
||||
default=InferenceArgs.output_path,
|
||||
default=FastVideoArgs.output_path,
|
||||
help="Directory to save generated videos",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--seed",
|
||||
type=int,
|
||||
default=InferenceArgs.seed,
|
||||
default=FastVideoArgs.seed,
|
||||
help="Random seed for reproducibility",
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "InferenceArgs":
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "FastVideoArgs":
|
||||
args.tp_size = args.tensor_parallel_size
|
||||
args.sp_size = args.sequence_parallel_size
|
||||
args.flow_shift = getattr(args, "shift", args.flow_shift)
|
||||
@@ -386,6 +433,15 @@ class InferenceArgs:
|
||||
|
||||
def check_inference_args(self) -> None:
|
||||
"""Validate inference arguments for consistency"""
|
||||
if self.tp_size is None:
|
||||
self.tp_size = self.num_gpus
|
||||
if self.sp_size is None:
|
||||
self.sp_size = self.num_gpus
|
||||
|
||||
if self.tp_size != self.sp_size:
|
||||
raise ValueError(
|
||||
f"tp_size ({self.tp_size}) must be equal to sp_size ({self.sp_size})"
|
||||
)
|
||||
|
||||
# Validate VAE spatial parallelism with VAE tiling
|
||||
if self.vae_sp and not self.vae_tiling:
|
||||
@@ -396,10 +452,10 @@ class InferenceArgs:
|
||||
raise ValueError("prompt_path must be a text file")
|
||||
|
||||
|
||||
_inference_args = None
|
||||
_current_fastvideo_args = None
|
||||
|
||||
|
||||
def prepare_inference_args(argv: List[str]) -> InferenceArgs:
|
||||
def prepare_fastvideo_args(argv: List[str]) -> FastVideoArgs:
|
||||
"""
|
||||
Prepare the inference arguments from the command line arguments.
|
||||
|
||||
@@ -411,26 +467,38 @@ def prepare_inference_args(argv: List[str]) -> InferenceArgs:
|
||||
The inference arguments.
|
||||
"""
|
||||
parser = FlexibleArgumentParser()
|
||||
InferenceArgs.add_cli_args(parser)
|
||||
FastVideoArgs.add_cli_args(parser)
|
||||
raw_args = parser.parse_args(argv)
|
||||
inference_args = InferenceArgs.from_cli_args(raw_args)
|
||||
inference_args.check_inference_args()
|
||||
global _inference_args
|
||||
_inference_args = inference_args
|
||||
return inference_args
|
||||
fastvideo_args = FastVideoArgs.from_cli_args(raw_args)
|
||||
fastvideo_args.check_inference_args()
|
||||
global _current_fastvideo_args
|
||||
_current_fastvideo_args = fastvideo_args
|
||||
return fastvideo_args
|
||||
|
||||
|
||||
def get_inference_args() -> InferenceArgs:
|
||||
global _inference_args
|
||||
if _inference_args is None:
|
||||
raise ValueError("Inference arguments not set")
|
||||
return _inference_args
|
||||
@contextmanager
|
||||
def set_current_fastvideo_args(fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Temporarily set the current fastvideo config.
|
||||
Used during model initialization.
|
||||
We save the current fastvideo config in a global variable,
|
||||
so that all modules can access it, e.g. custom ops
|
||||
can access the fastvideo config to determine how to dispatch.
|
||||
"""
|
||||
global _current_fastvideo_args
|
||||
old_fastvideo_args = _current_fastvideo_args
|
||||
try:
|
||||
_current_fastvideo_args = fastvideo_args
|
||||
yield
|
||||
finally:
|
||||
_current_fastvideo_args = old_fastvideo_args
|
||||
|
||||
|
||||
class DeprecatedAction(argparse.Action):
|
||||
|
||||
def __init__(self, option_strings, dest, nargs=0, **kwargs):
|
||||
super().__init__(option_strings, dest, nargs=nargs, **kwargs)
|
||||
|
||||
def __call__(self, parser, namespace, values, option_string=None):
|
||||
raise ValueError(self.help)
|
||||
def get_current_fastvideo_args() -> FastVideoArgs:
|
||||
if _current_fastvideo_args is None:
|
||||
# in ci, usually when we test custom ops/modules directly,
|
||||
# we don't set the fastvideo config. In that case, we set a default
|
||||
# config.
|
||||
# TODO(will): may need to handle this for CI.
|
||||
raise ValueError("Current fastvideo args is not set.")
|
||||
return _current_fastvideo_args
|
||||
@@ -9,7 +9,7 @@ from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -52,7 +52,7 @@ def get_forward_context() -> ForwardContext:
|
||||
@contextmanager
|
||||
def set_forward_context(current_timestep,
|
||||
attn_metadata,
|
||||
inference_args: InferenceArgs = None):
|
||||
fastvideo_args: Optional[FastVideoArgs] = None):
|
||||
"""A context manager that stores the current forward context,
|
||||
can be attention metadata, etc.
|
||||
Here we can inject common logic for every model forward pass.
|
||||
|
||||
@@ -10,7 +10,7 @@ from typing import Any, Dict
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines import (ComposedPipelineBase, ForwardBatch,
|
||||
build_pipeline)
|
||||
@@ -28,29 +28,29 @@ class InferenceEngine:
|
||||
def __init__(
|
||||
self,
|
||||
pipeline: ComposedPipelineBase,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
):
|
||||
"""
|
||||
Initialize the inference engine.
|
||||
|
||||
Args:
|
||||
pipeline: The pipeline to use for inference.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
default_negative_prompt: The default negative prompt to use.
|
||||
"""
|
||||
self.pipeline = pipeline
|
||||
self.inference_args = inference_args
|
||||
self.fastvideo_args = fastvideo_args
|
||||
|
||||
@classmethod
|
||||
def create_engine(
|
||||
cls,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> "InferenceEngine":
|
||||
"""
|
||||
Create an inference engine with the specified arguments.
|
||||
|
||||
Args:
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
model_loader_cls: The model loader class to use. If None, it will be
|
||||
determined from the model type.
|
||||
pipeline_type: The type of pipeline to create. If None, it will be
|
||||
@@ -71,16 +71,16 @@ class InferenceEngine:
|
||||
# this way for training we can just do pipeline_cls.from_pretrained(
|
||||
# checkpoint_path) and have it handle everything.
|
||||
# TODO(Peiyuan): Then maybe we should only pass in model path and device, not the entire inference args?
|
||||
pipeline = build_pipeline(inference_args)
|
||||
pipeline = build_pipeline(fastvideo_args)
|
||||
logger.info("Pipeline Ready")
|
||||
|
||||
# Create the inference engine
|
||||
return cls(pipeline, inference_args)
|
||||
return cls(pipeline, fastvideo_args)
|
||||
|
||||
def run(
|
||||
self,
|
||||
prompt: str,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Run inference with the pipeline.
|
||||
@@ -96,16 +96,17 @@ class InferenceEngine:
|
||||
"""
|
||||
out_dict: Dict[str, Any] = dict()
|
||||
|
||||
num_videos_per_prompt = inference_args.num_videos
|
||||
seed = inference_args.seed
|
||||
height = inference_args.height
|
||||
width = inference_args.width
|
||||
video_length = inference_args.num_frames
|
||||
negative_prompt = inference_args.neg_prompt
|
||||
infer_steps = inference_args.num_inference_steps
|
||||
guidance_scale = inference_args.guidance_scale
|
||||
flow_shift = inference_args.flow_shift
|
||||
embedded_guidance_scale = inference_args.embedded_cfg_scale
|
||||
num_videos_per_prompt = fastvideo_args.num_videos
|
||||
seed = fastvideo_args.seed
|
||||
height = fastvideo_args.height
|
||||
width = fastvideo_args.width
|
||||
video_length = fastvideo_args.num_frames
|
||||
negative_prompt = fastvideo_args.neg_prompt
|
||||
infer_steps = fastvideo_args.num_inference_steps
|
||||
guidance_scale = fastvideo_args.guidance_scale
|
||||
flow_shift = fastvideo_args.flow_shift
|
||||
embedded_guidance_scale = fastvideo_args.embedded_cfg_scale
|
||||
image_path = fastvideo_args.image_path
|
||||
|
||||
# ========================================================================
|
||||
# Arguments: target_width, target_height, target_video_length
|
||||
@@ -160,20 +161,21 @@ class InferenceEngine:
|
||||
# return
|
||||
# sp_group = get_sp_group()
|
||||
# local_rank = sp_group.rank
|
||||
device = torch.device(inference_args.device_str)
|
||||
device = torch.device(fastvideo_args.device_str)
|
||||
batch = ForwardBatch(
|
||||
image_path=image_path,
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
height=inference_args.height,
|
||||
width=inference_args.width,
|
||||
num_frames=inference_args.num_frames,
|
||||
num_inference_steps=inference_args.num_inference_steps,
|
||||
guidance_scale=inference_args.guidance_scale,
|
||||
height=fastvideo_args.height,
|
||||
width=fastvideo_args.width,
|
||||
num_frames=fastvideo_args.num_frames,
|
||||
num_inference_steps=fastvideo_args.num_inference_steps,
|
||||
guidance_scale=fastvideo_args.guidance_scale,
|
||||
# generator=generator,
|
||||
eta=0.0,
|
||||
n_tokens=n_tokens,
|
||||
data_type="video" if inference_args.num_frames > 1 else "image",
|
||||
data_type="video" if fastvideo_args.num_frames > 1 else "image",
|
||||
device=device,
|
||||
extra={}, # Any additional parameters
|
||||
)
|
||||
@@ -182,7 +184,7 @@ class InferenceEngine:
|
||||
print(batch)
|
||||
print('===============================================')
|
||||
print('===============================================')
|
||||
print(inference_args)
|
||||
print(fastvideo_args)
|
||||
|
||||
# ========================================================================
|
||||
# Pipeline inference
|
||||
@@ -190,7 +192,7 @@ class InferenceEngine:
|
||||
start_time = time.time()
|
||||
samples = self.pipeline.forward(
|
||||
batch=batch,
|
||||
inference_args=inference_args,
|
||||
fastvideo_args=fastvideo_args,
|
||||
).output
|
||||
# TODO(will): fix and move to hunyuan stage
|
||||
# out_dict["seeds"] = batch.seeds
|
||||
|
||||
@@ -23,7 +23,7 @@ class SiluAndMul(CustomOp):
|
||||
return: (num_tokens, d) or (batch_size, seq_len, d)
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
if current_platform.is_cuda_alike() or current_platform.is_cpu():
|
||||
self.op = torch.ops._C.silu_and_mul
|
||||
|
||||
@@ -113,7 +113,7 @@ class ScaleResidual(nn.Module):
|
||||
Applies gated residual connection.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
def __init__(self, prefix: str = ""):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, residual: torch.Tensor, x: torch.Tensor,
|
||||
@@ -139,6 +139,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
eps: float = 1e-6,
|
||||
elementwise_affine: bool = False,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
if norm_type == "rms":
|
||||
@@ -189,6 +190,7 @@ class LayerNormScaleShift(nn.Module):
|
||||
eps: float = 1e-6,
|
||||
elementwise_affine: bool = False,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
if norm_type == "rms":
|
||||
|
||||
@@ -778,12 +778,13 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
# no need to narrow
|
||||
is_sharded_weight = is_sharded_weight
|
||||
|
||||
shard_idx = 0
|
||||
param_data = param_data.narrow(output_dim, shard_offset, shard_size)
|
||||
if loaded_shard_id == "q":
|
||||
shard_id = tp_rank
|
||||
shard_idx = tp_rank
|
||||
else:
|
||||
shard_id = tp_rank // self.num_kv_head_replicas
|
||||
start_idx = shard_id * shard_size
|
||||
shard_idx = tp_rank // self.num_kv_head_replicas
|
||||
start_idx = shard_idx * shard_size
|
||||
|
||||
if not is_sharded_weight:
|
||||
loaded_weight = loaded_weight.narrow(output_dim, start_idx,
|
||||
|
||||
@@ -12,7 +12,6 @@ from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
class MLP(nn.Module):
|
||||
"""
|
||||
MLP for DiT blocks, NO gated linear units
|
||||
TODO: add Tensor Parallel
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
@@ -23,6 +22,7 @@ class MLP(nn.Module):
|
||||
bias: bool = True,
|
||||
act_type: str = "gelu_pytorch_tanh",
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.fc_in = ReplicatedLinear(
|
||||
|
||||
@@ -84,7 +84,7 @@ class RotaryEmbedding(CustomOp):
|
||||
head_size: int,
|
||||
rotary_dim: int,
|
||||
max_position_embeddings: int,
|
||||
base: int,
|
||||
base: Union[int, float],
|
||||
is_neox_style: bool,
|
||||
dtype: torch.dtype,
|
||||
) -> None:
|
||||
@@ -446,7 +446,7 @@ def get_rope(
|
||||
head_size: int,
|
||||
rotary_dim: int,
|
||||
max_position: int,
|
||||
base: int,
|
||||
base: Union[int, float],
|
||||
is_neox_style: bool = True,
|
||||
rope_scaling: Optional[Dict[str, Any]] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
|
||||
@@ -32,7 +32,8 @@ class PatchEmbed(nn.Module):
|
||||
norm_layer=None,
|
||||
flatten=True,
|
||||
bias=True,
|
||||
dtype=None):
|
||||
dtype=None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
# Convert patch_size to 2-tuple
|
||||
if isinstance(patch_size, (list, tuple)):
|
||||
@@ -73,6 +74,7 @@ class TimestepEmbedder(nn.Module):
|
||||
max_period=10000,
|
||||
dtype=None,
|
||||
freq_dtype=torch.float32,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
@@ -132,6 +134,7 @@ class ModulateProjection(nn.Module):
|
||||
factor: int = 2,
|
||||
act_layer: str = "silu",
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.factor = factor
|
||||
@@ -148,7 +151,7 @@ class ModulateProjection(nn.Module):
|
||||
return x
|
||||
|
||||
|
||||
def unpatchify(x, t, h, w, patch_size, channels):
|
||||
def unpatchify(x, t, h, w, patch_size, channels) -> torch.Tensor:
|
||||
"""
|
||||
Convert patched representation back to image space.
|
||||
|
||||
|
||||
@@ -1,15 +1,59 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import List, Optional, Union
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
# TODO
|
||||
class BaseDiT(nn.Module):
|
||||
class BaseDiT(nn.Module, ABC):
|
||||
_fsdp_shard_conditions: list = []
|
||||
attention_head_dim: int | None = None
|
||||
_param_names_mapping: dict
|
||||
hidden_size: int
|
||||
num_attention_heads: int
|
||||
_supported_attention_backends: List[_Backend] = []
|
||||
|
||||
def __init_subclass__(cls) -> None:
|
||||
required_class_attrs = [
|
||||
"_fsdp_shard_conditions", "_param_names_mapping"
|
||||
]
|
||||
super().__init_subclass__()
|
||||
for attr in required_class_attrs:
|
||||
if not hasattr(cls, attr):
|
||||
raise AttributeError(
|
||||
f"Subclasses of BaseDiT must define '{attr}' class variable"
|
||||
)
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
super().__init__()
|
||||
if not self.supported_attention_backends:
|
||||
raise ValueError(
|
||||
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
|
||||
)
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
@abstractmethod
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states_image: Optional[Union[
|
||||
torch.Tensor, List[torch.Tensor]]] = None,
|
||||
guidance=None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
pass
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
required_attrs = ["hidden_size", "num_attention_heads"]
|
||||
for attr in required_attrs:
|
||||
if not hasattr(self, attr):
|
||||
raise AttributeError(
|
||||
f"Subclasses of BaseDiT must define '{attr}' instance variable"
|
||||
)
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> List[_Backend]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
@@ -19,6 +19,7 @@ from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
|
||||
PatchEmbed, TimestepEmbedder,
|
||||
unpatchify)
|
||||
from fastvideo.v1.models.dits.base import BaseDiT
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
class HunyuanRMSNorm(nn.Module):
|
||||
@@ -91,6 +92,8 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
supported_attention_backends: Optional[List[_Backend]] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -105,6 +108,7 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
factor=6,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.img_mod",
|
||||
)
|
||||
|
||||
# Fused operations for image stream
|
||||
@@ -123,7 +127,8 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
self.img_attn_qkv = ReplicatedLinear(hidden_size,
|
||||
hidden_size * 3,
|
||||
bias=True,
|
||||
params_dtype=dtype)
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.img_attn_qkv")
|
||||
|
||||
self.img_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
self.img_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
@@ -131,9 +136,14 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
self.img_attn_proj = ReplicatedLinear(hidden_size,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype)
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.img_attn_proj")
|
||||
|
||||
self.img_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
|
||||
self.img_mlp = MLP(hidden_size,
|
||||
mlp_hidden_dim,
|
||||
bias=True,
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.img_mlp")
|
||||
|
||||
# Text modulation components
|
||||
self.txt_mod = ModulateProjection(
|
||||
@@ -141,6 +151,7 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
factor=6,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.txt_mod",
|
||||
)
|
||||
|
||||
# Fused operations for text stream
|
||||
@@ -173,27 +184,12 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
|
||||
|
||||
# Distributed attention
|
||||
self.attn = DistributedAttention(num_heads=num_attention_heads,
|
||||
head_size=head_dim,
|
||||
dropout_rate=0.0,
|
||||
causal=False)
|
||||
|
||||
# QK norm layers for text
|
||||
self.txt_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
self.txt_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
|
||||
self.txt_attn_proj = ReplicatedLinear(hidden_size,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype)
|
||||
|
||||
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
|
||||
|
||||
# Distributed attention
|
||||
self.attn = DistributedAttention(num_heads=num_attention_heads,
|
||||
head_size=head_dim,
|
||||
dropout_rate=0.0,
|
||||
causal=False)
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_attention_heads,
|
||||
head_size=head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix=f"{prefix}.attn")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -303,6 +299,8 @@ class MMSingleStreamBlock(nn.Module):
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
supported_attention_backends: Optional[List[_Backend]] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -317,13 +315,15 @@ class MMSingleStreamBlock(nn.Module):
|
||||
self.linear1 = ReplicatedLinear(hidden_size,
|
||||
hidden_size * 3 + mlp_hidden_dim,
|
||||
bias=True,
|
||||
params_dtype=dtype)
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.linear1")
|
||||
|
||||
# Combined projection and MLP output
|
||||
self.linear2 = ReplicatedLinear(hidden_size + mlp_hidden_dim,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype)
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.linear2")
|
||||
|
||||
# QK norm layers
|
||||
self.q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
|
||||
@@ -345,13 +345,16 @@ class MMSingleStreamBlock(nn.Module):
|
||||
self.modulation = ModulateProjection(hidden_size,
|
||||
factor=3,
|
||||
act_layer="silu",
|
||||
dtype=dtype)
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.modulation")
|
||||
|
||||
# Distributed attention
|
||||
self.attn = DistributedAttention(num_heads=num_attention_heads,
|
||||
head_size=head_dim,
|
||||
dropout_rate=0.0,
|
||||
causal=False)
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_attention_heads,
|
||||
head_size=head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix=f"{prefix}.attn")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -433,6 +436,9 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
|
||||
lambda n, m: "single" in n and str.isdigit(n.split(".")[-1]),
|
||||
lambda n, m: "refiner" in n and str.isdigit(n.split(".")[-1]),
|
||||
]
|
||||
_supported_attention_backends = [
|
||||
_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN, _Backend.TORCH_SDPA
|
||||
]
|
||||
_param_names_mapping = {
|
||||
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
|
||||
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
|
||||
@@ -548,24 +554,25 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
|
||||
}
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: int = 2,
|
||||
patch_size_t: int = 1,
|
||||
in_channels: int = 16,
|
||||
out_channels: int = 16,
|
||||
num_attention_heads: int = 24,
|
||||
attention_head_dim: int = 128,
|
||||
mlp_ratio: float = 4.0,
|
||||
num_layers: int = 20,
|
||||
num_single_layers: int = 40,
|
||||
num_refiner_layers: int = 2,
|
||||
rope_axes_dim: Tuple[int, int, int] = (16, 56, 56),
|
||||
guidance_embeds: bool = False,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
text_embed_dim: int = 4096,
|
||||
pooled_projection_dim: int = 768,
|
||||
rope_theta: int = 256,
|
||||
qk_norm: str = "rms_norm", #TODO(PY)
|
||||
self,
|
||||
patch_size: int = 2,
|
||||
patch_size_t: int = 1,
|
||||
in_channels: int = 16,
|
||||
out_channels: int = 16,
|
||||
num_attention_heads: int = 24,
|
||||
attention_head_dim: int = 128,
|
||||
mlp_ratio: float = 4.0,
|
||||
num_layers: int = 20,
|
||||
num_single_layers: int = 40,
|
||||
num_refiner_layers: int = 2,
|
||||
rope_axes_dim: Tuple[int, int, int] = (16, 56, 56),
|
||||
guidance_embeds: bool = False,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
text_embed_dim: int = 4096,
|
||||
pooled_projection_dim: int = 768,
|
||||
rope_theta: int = 256,
|
||||
qk_norm: str = "rms_norm", #TODO(PY)
|
||||
prefix="",
|
||||
):
|
||||
super().__init__()
|
||||
hidden_size = attention_head_dim * num_attention_heads
|
||||
@@ -598,29 +605,35 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
|
||||
self.img_in = PatchEmbed(self.patch_size,
|
||||
self.in_channels,
|
||||
self.hidden_size,
|
||||
dtype=dtype)
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.img_in")
|
||||
|
||||
self.txt_in = SingleTokenRefiner(self.text_states_dim,
|
||||
hidden_size,
|
||||
num_attention_heads,
|
||||
depth=num_refiner_layers,
|
||||
dtype=dtype)
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.txt_in")
|
||||
|
||||
# Time modulation
|
||||
self.time_in = TimestepEmbedder(self.hidden_size,
|
||||
act_layer="silu",
|
||||
dtype=dtype)
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.time_in")
|
||||
|
||||
# Text modulation
|
||||
self.vector_in = MLP(self.text_states_dim_2,
|
||||
self.hidden_size,
|
||||
self.hidden_size,
|
||||
act_type="silu",
|
||||
dtype=dtype)
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.vector_in")
|
||||
|
||||
# Guidance modulation
|
||||
self.guidance_in = (TimestepEmbedder(
|
||||
self.hidden_size, act_layer="silu", dtype=dtype)
|
||||
self.guidance_in = (TimestepEmbedder(self.hidden_size,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.guidance_in")
|
||||
if self.guidance_embeds else None)
|
||||
|
||||
# Double blocks
|
||||
@@ -630,7 +643,8 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
|
||||
num_attention_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
dtype=dtype,
|
||||
) for _ in range(num_layers)
|
||||
supported_attention_backends=self._supported_attention_backends,
|
||||
prefix=f"{prefix}.double_blocks.{i}") for i in range(num_layers)
|
||||
])
|
||||
|
||||
# Single blocks
|
||||
@@ -640,23 +654,29 @@ class HunyuanVideoTransformer3DModel(BaseDiT):
|
||||
num_attention_heads,
|
||||
mlp_ratio=mlp_ratio,
|
||||
dtype=dtype,
|
||||
) for _ in range(num_single_layers)
|
||||
supported_attention_backends=self._supported_attention_backends,
|
||||
prefix=f"{prefix}.single_blocks.{i+num_layers}")
|
||||
for i in range(num_single_layers)
|
||||
])
|
||||
|
||||
self.final_layer = FinalLayer(hidden_size,
|
||||
self.patch_size,
|
||||
self.out_channels,
|
||||
dtype=dtype)
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.final_layer")
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
# TODO: change the input the FORWAD_BACTCH Dict
|
||||
# TODO: change output to a dict
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
|
||||
timestep: torch.LongTensor,
|
||||
guidance=None,
|
||||
):
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states_image: Optional[Union[
|
||||
torch.Tensor, List[torch.Tensor]]] = None,
|
||||
guidance=None,
|
||||
**kwargs):
|
||||
"""
|
||||
Forward pass of the HunyuanDiT model.
|
||||
|
||||
@@ -758,26 +778,31 @@ class SingleTokenRefiner(nn.Module):
|
||||
depth=2,
|
||||
qkv_bias=True,
|
||||
dtype=None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
# Input projection
|
||||
self.input_embedder = ReplicatedLinear(in_channels,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype)
|
||||
self.input_embedder = ReplicatedLinear(
|
||||
in_channels,
|
||||
hidden_size,
|
||||
bias=True,
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.input_embedder")
|
||||
|
||||
# Timestep embedding
|
||||
self.t_embedder = TimestepEmbedder(hidden_size,
|
||||
act_layer="silu",
|
||||
dtype=dtype)
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.t_embedder")
|
||||
|
||||
# Context embedding
|
||||
self.c_embedder = MLP(in_channels,
|
||||
hidden_size,
|
||||
hidden_size,
|
||||
act_type="silu",
|
||||
dtype=dtype)
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.c_embedder")
|
||||
|
||||
# Refiner blocks
|
||||
self.refiner_blocks = nn.ModuleList([
|
||||
@@ -786,7 +811,8 @@ class SingleTokenRefiner(nn.Module):
|
||||
num_attention_heads,
|
||||
qkv_bias=qkv_bias,
|
||||
dtype=dtype,
|
||||
) for _ in range(depth)
|
||||
prefix=f"{prefix}.refiner_blocks.{i}",
|
||||
) for i in range(depth)
|
||||
])
|
||||
|
||||
def forward(self, x, t):
|
||||
@@ -820,6 +846,7 @@ class IndividualTokenRefinerBlock(nn.Module):
|
||||
mlp_ratio=4.0,
|
||||
qkv_bias=True,
|
||||
dtype=None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.num_attention_heads = num_attention_heads
|
||||
@@ -834,12 +861,15 @@ class IndividualTokenRefinerBlock(nn.Module):
|
||||
self.self_attn_qkv = ReplicatedLinear(hidden_size,
|
||||
hidden_size * 3,
|
||||
bias=qkv_bias,
|
||||
params_dtype=dtype)
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.self_attn_qkv")
|
||||
|
||||
self.self_attn_proj = ReplicatedLinear(hidden_size,
|
||||
hidden_size,
|
||||
bias=qkv_bias,
|
||||
params_dtype=dtype)
|
||||
self.self_attn_proj = ReplicatedLinear(
|
||||
hidden_size,
|
||||
hidden_size,
|
||||
bias=qkv_bias,
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.self_attn_proj")
|
||||
|
||||
# MLP
|
||||
self.norm2 = nn.LayerNorm(hidden_size,
|
||||
@@ -850,18 +880,25 @@ class IndividualTokenRefinerBlock(nn.Module):
|
||||
mlp_hidden_dim,
|
||||
bias=True,
|
||||
act_type="silu",
|
||||
dtype=dtype)
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.mlp")
|
||||
|
||||
# Modulation
|
||||
self.adaLN_modulation = ModulateProjection(hidden_size,
|
||||
factor=2,
|
||||
act_layer="silu",
|
||||
dtype=dtype)
|
||||
self.adaLN_modulation = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=2,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.adaLN_modulation")
|
||||
|
||||
# Scaled dot product attention
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_attention_heads,
|
||||
head_size=hidden_size // num_attention_heads,
|
||||
# TODO: remove hardcode; remove STA
|
||||
supported_attention_backends=[
|
||||
_Backend.FLASH_ATTN, _Backend.TORCH_SDPA
|
||||
],
|
||||
)
|
||||
|
||||
def forward(self, x, c):
|
||||
@@ -900,7 +937,8 @@ class FinalLayer(nn.Module):
|
||||
hidden_size,
|
||||
patch_size,
|
||||
out_channels,
|
||||
dtype=None) -> None:
|
||||
dtype=None,
|
||||
prefix: str = "") -> None:
|
||||
super().__init__()
|
||||
|
||||
# Normalization
|
||||
@@ -914,13 +952,16 @@ class FinalLayer(nn.Module):
|
||||
self.linear = ReplicatedLinear(hidden_size,
|
||||
output_dim,
|
||||
bias=True,
|
||||
params_dtype=dtype)
|
||||
params_dtype=dtype,
|
||||
prefix=f"{prefix}.linear")
|
||||
|
||||
# Modulation
|
||||
self.adaLN_modulation = ModulateProjection(hidden_size,
|
||||
factor=2,
|
||||
act_layer="silu",
|
||||
dtype=dtype)
|
||||
self.adaLN_modulation = ModulateProjection(
|
||||
hidden_size,
|
||||
factor=2,
|
||||
act_layer="silu",
|
||||
dtype=dtype,
|
||||
prefix=f"{prefix}.adaLN_modulation")
|
||||
|
||||
def forward(self, x, c):
|
||||
# What the heck HF? Why you change the scale and shift order here???
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from typing import Any, Dict, Optional, Tuple, Union
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
@@ -21,6 +21,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
|
||||
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
|
||||
PatchEmbed, TimestepEmbedder)
|
||||
from fastvideo.v1.models.dits.base import BaseDiT
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
class WanImageEmbedding(torch.nn.Module):
|
||||
@@ -34,9 +35,10 @@ class WanImageEmbedding(torch.nn.Module):
|
||||
|
||||
def forward(self,
|
||||
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
dtype = encoder_hidden_states_image.dtype
|
||||
hidden_states = self.norm1(encoder_hidden_states_image)
|
||||
hidden_states = self.ff(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states).to(dtype)
|
||||
return hidden_states
|
||||
|
||||
|
||||
@@ -52,10 +54,7 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
self.time_embedder = TimestepEmbedder(
|
||||
dim,
|
||||
frequency_embedding_size=time_freq_dim,
|
||||
act_layer="silu",
|
||||
freq_dtype=torch.float64)
|
||||
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
|
||||
self.time_modulation = ModulateProjection(dim,
|
||||
factor=6,
|
||||
act_layer="silu")
|
||||
@@ -75,9 +74,8 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_hidden_states_image: Optional[torch.Tensor] = None,
|
||||
):
|
||||
with torch.cuda.amp.autocast(dtype=torch.float32):
|
||||
temb = self.time_embedder(timestep.float())
|
||||
timestep_proj = self.time_modulation(temb)
|
||||
temb = self.time_embedder(timestep)
|
||||
timestep_proj = self.time_modulation(temb)
|
||||
|
||||
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
|
||||
if encoder_hidden_states_image is not None:
|
||||
@@ -116,9 +114,14 @@ class WanSelfAttention(nn.Module):
|
||||
self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
# Scaled dot product attention
|
||||
self.attn = LocalAttention(dropout_rate=0,
|
||||
self.attn = LocalAttention(num_heads=num_heads,
|
||||
head_size=self.head_dim,
|
||||
dropout_rate=0,
|
||||
softmax_scale=None,
|
||||
causal=False)
|
||||
causal=False,
|
||||
supported_attention_backends=[
|
||||
_Backend.FLASH_ATTN, _Backend.TORCH_SDPA
|
||||
])
|
||||
|
||||
def forward(self, x: torch.Tensor, context: torch.Tensor,
|
||||
context_lens: int):
|
||||
@@ -144,8 +147,8 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
v = self.to_v(context)[0].view(b, -1, n, d)
|
||||
|
||||
# compute attention
|
||||
@@ -159,13 +162,16 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
|
||||
class WanI2VCrossAttention(WanSelfAttention):
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
num_heads: int,
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
eps=1e-6) -> None:
|
||||
super().__init__(dim, num_heads, window_size, qk_norm, eps)
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
num_heads: int,
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
eps=1e-6,
|
||||
supported_attention_backends: Optional[List[str]] = None) -> None:
|
||||
super().__init__(dim, num_heads, window_size, qk_norm, eps,
|
||||
supported_attention_backends)
|
||||
|
||||
self.add_k_proj = ReplicatedLinear(dim, dim)
|
||||
self.add_v_proj = ReplicatedLinear(dim, dim)
|
||||
@@ -184,11 +190,11 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
b, n, d = x.size(0), self.num_heads, self.head_dim
|
||||
|
||||
# compute query, key, value
|
||||
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
v = self.to_v(context)[0].view(b, -1, n, d)
|
||||
k_img = self.norm_added_k.forward_native(
|
||||
self.add_k_proj(context_img)[0]).view(b, -1, n, d)
|
||||
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
|
||||
b, -1, n, d)
|
||||
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
|
||||
img_x = self.attn(q, k_img, v_img)
|
||||
# compute attention
|
||||
@@ -213,6 +219,7 @@ class WanTransformerBlock(nn.Module):
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
supported_attention_backends: Optional[List[_Backend]] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
@@ -222,10 +229,11 @@ class WanTransformerBlock(nn.Module):
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_out = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.attn1 = DistributedAttention(num_heads=num_heads,
|
||||
head_size=dim // num_heads,
|
||||
dropout_rate=0.0,
|
||||
causal=False)
|
||||
self.attn1 = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=dim // num_heads,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends)
|
||||
self.hidden_dim = dim
|
||||
self.num_attention_heads = num_heads
|
||||
dim_head = dim // num_heads
|
||||
@@ -284,16 +292,15 @@ class WanTransformerBlock(nn.Module):
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
assert temb.dtype == torch.float32
|
||||
with torch.cuda.amp.autocast(dtype=torch.float32):
|
||||
e = self.scale_shift_table + temb
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb.float()
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = self.norm1(hidden_states.float()).to(
|
||||
dtype=orig_dtype) * (1 + scale_msa) + shift_msa
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) *
|
||||
(1 + scale_msa) + shift_msa).to(orig_dtype)
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
@@ -320,6 +327,8 @@ class WanTransformerBlock(nn.Module):
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
@@ -327,10 +336,13 @@ class WanTransformerBlock(nn.Module):
|
||||
context_lens=None)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
@@ -339,6 +351,7 @@ class WanTransformer3DModel(BaseDiT):
|
||||
_fsdp_shard_conditions = [
|
||||
lambda n, m: "blocks" in n and str.isdigit(n.split(".")[-1]),
|
||||
]
|
||||
_supported_attention_backends = [_Backend.FLASH_ATTN, _Backend.TORCH_SDPA]
|
||||
_param_names_mapping = {
|
||||
r"^patch_embedding\.(.*)$":
|
||||
r"patch_embedding.proj.\1",
|
||||
@@ -400,8 +413,9 @@ class WanTransformer3DModel(BaseDiT):
|
||||
super().__init__()
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
self.inner_dim = inner_dim
|
||||
self.hidden_size = inner_dim
|
||||
self.num_attention_heads = num_attention_heads
|
||||
self.in_channels = in_channels
|
||||
self.out_channels = out_channels or in_channels
|
||||
self.patch_size = patch_size
|
||||
self.text_len = text_len
|
||||
@@ -424,7 +438,9 @@ class WanTransformer3DModel(BaseDiT):
|
||||
self.blocks = nn.ModuleList([
|
||||
WanTransformerBlock(inner_dim, ffn_dim, num_attention_heads,
|
||||
qk_norm, cross_attn_norm, eps,
|
||||
added_kv_proj_dim) for _ in range(num_layers)
|
||||
added_kv_proj_dim,
|
||||
self._supported_attention_backends)
|
||||
for _ in range(num_layers)
|
||||
])
|
||||
|
||||
# 4. Output norm & projection
|
||||
@@ -440,19 +456,24 @@ class WanTransformer3DModel(BaseDiT):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
seq_len: Optional[int] = None,
|
||||
encoder_hidden_states_image: Optional[torch.Tensor] = None,
|
||||
y: Optional[torch.Tensor] = None,
|
||||
return_dict: bool = True,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||
if y is not None:
|
||||
hidden_states = torch.cat([hidden_states, y], dim=1)
|
||||
self.__post_init__()
|
||||
|
||||
def forward(self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Union[torch.Tensor, List[torch.Tensor]],
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states_image: Optional[Union[
|
||||
torch.Tensor, List[torch.Tensor]]] = None,
|
||||
guidance=None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
orig_dtype = hidden_states.dtype
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
if isinstance(encoder_hidden_states_image,
|
||||
list) and len(encoder_hidden_states_image) > 0:
|
||||
encoder_hidden_states_image = encoder_hidden_states_image[0]
|
||||
else:
|
||||
encoder_hidden_states_image = None
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.patch_size
|
||||
@@ -461,40 +482,23 @@ class WanTransformer3DModel(BaseDiT):
|
||||
post_patch_width = width // p_w
|
||||
|
||||
# Get rotary embeddings
|
||||
d = self.inner_dim // self.num_attention_heads
|
||||
d = self.hidden_size // self.num_attention_heads
|
||||
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
|
||||
freqs_cos, freqs_sin = get_rotary_pos_embed(
|
||||
(post_patch_num_frames * get_sequence_model_parallel_world_size(),
|
||||
post_patch_height, post_patch_width),
|
||||
self.inner_dim,
|
||||
self.hidden_size,
|
||||
self.num_attention_heads,
|
||||
rope_dim_list,
|
||||
dtype=torch.float64,
|
||||
rope_theta=10000)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
|
||||
freqs_cis = (freqs_cos.float(),
|
||||
freqs_sin.float()) if freqs_cos is not None else None
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
grid_sizes = torch.stack(
|
||||
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
if seq_len is None:
|
||||
seq_len = hidden_states.size(1)
|
||||
hidden_states = torch.cat([
|
||||
hidden_states,
|
||||
hidden_states.new_zeros(1, seq_len - hidden_states.size(1),
|
||||
hidden_states.size(2))
|
||||
],
|
||||
dim=1)
|
||||
|
||||
encoder_hidden_states = torch.cat([
|
||||
encoder_hidden_states,
|
||||
encoder_hidden_states.new_zeros(
|
||||
1, self.text_len - encoder_hidden_states.size(1),
|
||||
encoder_hidden_states.size(2))
|
||||
],
|
||||
dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image)
|
||||
@@ -504,6 +508,7 @@ class WanTransformer3DModel(BaseDiT):
|
||||
encoder_hidden_states = torch.concat(
|
||||
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
|
||||
assert encoder_hidden_states.dtype == orig_dtype
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.blocks:
|
||||
@@ -516,39 +521,16 @@ class WanTransformer3DModel(BaseDiT):
|
||||
timestep_proj, freqs_cis)
|
||||
|
||||
# 5. Output norm, projection & unpatchify
|
||||
with torch.cuda.amp.autocast(dtype=torch.float32):
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(
|
||||
2, dim=1)
|
||||
hidden_states = self.norm_out(hidden_states.float(), shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
|
||||
dim=1)
|
||||
hidden_states = self.norm_out(hidden_states.float(), shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
post_patch_height,
|
||||
post_patch_width, p_t, p_h, p_w,
|
||||
-1)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
return output.float()
|
||||
|
||||
def unpatchify(self, x, grid_sizes) -> torch.Tensor:
|
||||
r"""
|
||||
Reconstruct video tensors from patch embeddings.
|
||||
|
||||
Args:
|
||||
x (List[Tensor]):
|
||||
List of patchified features, each with shape [L, C_out * prod(patch_size)]
|
||||
grid_sizes (Tensor):
|
||||
Original spatial-temporal grid dimensions before patching,
|
||||
shape [B, 3] (3 dimensions correspond to F_patches, H_patches, W_patches)
|
||||
|
||||
Returns:
|
||||
Tensor:
|
||||
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
|
||||
"""
|
||||
|
||||
c = self.out_channels
|
||||
out = []
|
||||
for u, v in zip(x, grid_sizes.tolist()):
|
||||
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
|
||||
u = u.permute(6, 0, 3, 1, 4, 2, 5)
|
||||
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
out.append(u)
|
||||
out = torch.cat(out, dim=0)
|
||||
return out
|
||||
return output
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
from typing import List
|
||||
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
class BaseEncoder(nn.Module):
|
||||
_supported_attention_backends: List[_Backend] = []
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
super().__init__()
|
||||
if not self.supported_attention_backends:
|
||||
raise ValueError(
|
||||
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
|
||||
)
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
pass
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> List[_Backend]:
|
||||
return self._supported_attention_backends
|
||||
@@ -19,11 +19,13 @@ from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
|
||||
RowParallelLinear)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.encoders.base import BaseEncoder
|
||||
from fastvideo.v1.models.encoders.vision import (VisionEncoderInfo,
|
||||
resolve_visual_encoder_outputs)
|
||||
# TODO: support quantization
|
||||
# from vllm.model_executor.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -195,7 +197,9 @@ class CLIPAttention(nn.Module):
|
||||
self.head_dim,
|
||||
self.num_heads_per_partition,
|
||||
softmax_scale=self.scale,
|
||||
causal=True)
|
||||
causal=True,
|
||||
supported_attention_backends=self.config.
|
||||
supported_attention_backends)
|
||||
|
||||
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
|
||||
return tensor.view(bsz, seq_len, self.num_heads,
|
||||
@@ -464,7 +468,8 @@ class CLIPTextTransformer(nn.Module):
|
||||
)
|
||||
|
||||
|
||||
class CLIPTextModel(nn.Module):
|
||||
class CLIPTextModel(BaseEncoder):
|
||||
_supported_attention_backends = [_Backend.FLASH_ATTN, _Backend.TORCH_SDPA]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -475,6 +480,7 @@ class CLIPTextModel(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
self.config = config
|
||||
self.config.supported_attention_backends = self._supported_attention_backends
|
||||
self.text_model = CLIPTextTransformer(config=config,
|
||||
quant_config=quant_config,
|
||||
prefix=prefix)
|
||||
@@ -598,6 +604,9 @@ class CLIPVisionTransformer(nn.Module):
|
||||
inputs_embeds=hidden_states,
|
||||
return_all_hidden_states=return_all_hidden_states)
|
||||
|
||||
if not return_all_hidden_states:
|
||||
encoder_outputs = encoder_outputs[0]
|
||||
|
||||
# Handle post-norm (if applicable) and stacks feature layers if needed
|
||||
encoder_outputs = resolve_visual_encoder_outputs(
|
||||
encoder_outputs, feature_sample_layers, self.post_layernorm,
|
||||
@@ -606,10 +615,11 @@ class CLIPVisionTransformer(nn.Module):
|
||||
return encoder_outputs
|
||||
|
||||
|
||||
class CLIPVisionModel(nn.Module, SupportsQuant):
|
||||
class CLIPVisionModel(BaseEncoder, SupportsQuant):
|
||||
config_class = CLIPVisionConfig
|
||||
main_input_name = "pixel_values"
|
||||
packed_modules_mapping = {"qkv_proj": ["q_proj", "k_proj", "v_proj"]}
|
||||
_supported_attention_backends = [_Backend.FLASH_ATTN, _Backend.TORCH_SDPA]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -621,6 +631,8 @@ class CLIPVisionModel(nn.Module, SupportsQuant):
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.config.supported_attention_backends = self._supported_attention_backends
|
||||
self.vision_model = CLIPVisionTransformer(
|
||||
config=config,
|
||||
quant_config=quant_config,
|
||||
@@ -654,6 +666,8 @@ class CLIPVisionModel(nn.Module, SupportsQuant):
|
||||
layer_count = len(self.vision_model.encoder.layers)
|
||||
|
||||
for name, loaded_weight in weights:
|
||||
if name.startswith("visual_projection"):
|
||||
continue
|
||||
# post_layernorm is not needed in CLIPVisionModel
|
||||
if (name.startswith("vision_model.post_layernorm")
|
||||
and self.vision_model.post_layernorm is None):
|
||||
|
||||
@@ -39,10 +39,11 @@ from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
|
||||
QKVParallelLinear, RowParallelLinear)
|
||||
from fastvideo.v1.layers.rotary_embedding import get_rope
|
||||
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.v1.models.encoders.base import BaseEncoder
|
||||
from fastvideo.v1.models.loader.weight_utils import (default_weight_loader,
|
||||
maybe_remap_kv_scale_name)
|
||||
|
||||
# from ..utils import (extract_layer_index)
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
|
||||
|
||||
class QuantizationConfig:
|
||||
@@ -159,16 +160,18 @@ class LlamaAttention(nn.Module):
|
||||
self.head_dim,
|
||||
rotary_dim=self.rotary_dim,
|
||||
max_position=max_position_embeddings,
|
||||
base=rope_theta,
|
||||
base=int(rope_theta),
|
||||
rope_scaling=rope_scaling,
|
||||
is_neox_style=is_neox_style,
|
||||
)
|
||||
|
||||
self.attn = LocalAttention(self.num_heads,
|
||||
self.head_dim,
|
||||
self.num_kv_heads,
|
||||
softmax_scale=self.scaling,
|
||||
causal=True)
|
||||
self.attn = LocalAttention(
|
||||
self.num_heads,
|
||||
self.head_dim,
|
||||
self.num_kv_heads,
|
||||
softmax_scale=self.scaling,
|
||||
causal=True,
|
||||
supported_attention_backends=config.supported_attention_backends)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -276,7 +279,8 @@ class LlamaDecoderLayer(nn.Module):
|
||||
return hidden_states, residual
|
||||
|
||||
|
||||
class LlamaModel(nn.Module):
|
||||
class LlamaModel(BaseEncoder):
|
||||
_supported_attention_backends = [_Backend.FLASH_ATTN, _Backend.TORCH_SDPA]
|
||||
|
||||
def __init__(self,
|
||||
config: LlamaConfig,
|
||||
@@ -288,6 +292,7 @@ class LlamaModel(nn.Module):
|
||||
lora_config = None
|
||||
|
||||
self.config = config
|
||||
self.config.supported_attention_backends = self._supported_attention_backends
|
||||
self.quant_config = quant_config
|
||||
if lora_config is not None:
|
||||
max_loras = 1
|
||||
|
||||
@@ -28,11 +28,12 @@ import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from transformers import T5Config
|
||||
|
||||
from fastvideo.v1.distributed import get_tensor_model_parallel_world_size
|
||||
from fastvideo.v1.distributed import (get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size)
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.layers.layernorm import RMSNorm
|
||||
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
|
||||
RowParallelLinear)
|
||||
from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
|
||||
QKVParallelLinear, RowParallelLinear)
|
||||
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
|
||||
|
||||
@@ -67,7 +68,8 @@ class T5DenseActDense(nn.Module):
|
||||
config: T5Config,
|
||||
quant_config: Optional[QuantizationConfig] = None):
|
||||
super().__init__()
|
||||
self.wi = ColumnParallelLinear(config.d_model, config.d_ff, bias=False)
|
||||
self.wi = MergedColumnParallelLinear(config.d_model, [config.d_ff],
|
||||
bias=False)
|
||||
self.wo = RowParallelLinear(config.d_ff,
|
||||
config.d_model,
|
||||
bias=False,
|
||||
@@ -87,14 +89,12 @@ class T5DenseGatedActDense(nn.Module):
|
||||
config: T5Config,
|
||||
quant_config: Optional[QuantizationConfig] = None):
|
||||
super().__init__()
|
||||
self.wi_0 = ColumnParallelLinear(config.d_model,
|
||||
config.d_ff,
|
||||
bias=False,
|
||||
quant_config=quant_config)
|
||||
self.wi_1 = ColumnParallelLinear(config.d_model,
|
||||
config.d_ff,
|
||||
bias=False,
|
||||
quant_config=quant_config)
|
||||
self.wi_0 = MergedColumnParallelLinear(config.d_model, [config.d_ff],
|
||||
bias=False,
|
||||
quant_config=quant_config)
|
||||
self.wi_1 = MergedColumnParallelLinear(config.d_model, [config.d_ff],
|
||||
bias=False,
|
||||
quant_config=quant_config)
|
||||
# Should not run in fp16 unless mixed-precision is used,
|
||||
# see https://github.com/huggingface/transformers/issues/20287.
|
||||
self.wo = RowParallelLinear(config.d_ff,
|
||||
@@ -170,6 +170,7 @@ class T5Attention(nn.Module):
|
||||
config.relative_attention_max_distance
|
||||
self.d_model = config.d_model
|
||||
self.key_value_proj_dim = config.d_kv
|
||||
self.total_num_heads = self.total_num_kv_heads = config.num_heads
|
||||
|
||||
# Partition heads across multiple tensor parallel GPUs.
|
||||
tp_world_size = get_tensor_model_parallel_world_size()
|
||||
@@ -178,29 +179,33 @@ class T5Attention(nn.Module):
|
||||
|
||||
self.inner_dim = self.n_heads * self.key_value_proj_dim
|
||||
# No GQA in t5.
|
||||
self.n_kv_heads = self.n_heads
|
||||
# self.n_kv_heads = self.n_heads
|
||||
|
||||
self.qkv_proj = QKVParallelLinear(self.d_model,
|
||||
self.d_model // self.n_heads,
|
||||
self.n_heads,
|
||||
self.n_kv_heads,
|
||||
bias=False,
|
||||
quant_config=quant_config)
|
||||
self.qkv_proj = QKVParallelLinear(
|
||||
self.d_model,
|
||||
self.d_model // self.total_num_heads,
|
||||
self.total_num_heads,
|
||||
self.total_num_kv_heads,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.qkv_proj",
|
||||
)
|
||||
|
||||
self.attn = T5MultiHeadAttention()
|
||||
|
||||
if self.has_relative_attention_bias:
|
||||
self.relative_attention_bias = \
|
||||
VocabParallelEmbedding(self.relative_attention_num_buckets,
|
||||
self.n_heads,
|
||||
self.total_num_heads,
|
||||
org_num_embeddings=self.relative_attention_num_buckets,
|
||||
padding_size=self.relative_attention_num_buckets,
|
||||
quant_config=quant_config)
|
||||
self.o = RowParallelLinear(
|
||||
self.inner_dim,
|
||||
self.d_model,
|
||||
self.d_model,
|
||||
bias=False,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.o_proj",
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
@@ -295,13 +300,13 @@ class T5Attention(nn.Module):
|
||||
) -> torch.Tensor:
|
||||
bs, seq_len, _ = hidden_states.shape
|
||||
num_seqs = bs
|
||||
n, c = self.n_heads, self.d_model // self.n_heads
|
||||
n, c = self.n_heads, self.d_model // self.total_num_heads
|
||||
qkv, _ = self.qkv_proj(hidden_states)
|
||||
# Projection of 'own' hidden state (self-attention). No GQA here.
|
||||
q, k, v = qkv.split(self.inner_dim, dim=-1)
|
||||
q = q.view(bs, -1, n, c)
|
||||
k = k.view(bs, -1, n, c)
|
||||
v = v.view(bs, -1, n, c)
|
||||
q = q.reshape(bs, seq_len, n, c)
|
||||
k = k.reshape(bs, seq_len, n, c)
|
||||
v = v.reshape(bs, seq_len, n, c)
|
||||
|
||||
assert attn_metadata is not None
|
||||
attn_bias = attn_metadata.attn_bias
|
||||
@@ -325,6 +330,11 @@ class T5Attention(nn.Module):
|
||||
-1) if attention_mask.ndim == 2 else attention_mask.unsqueeze(1)
|
||||
attn_bias.masked_fill_(attention_mask == 0,
|
||||
torch.finfo(q.dtype).min)
|
||||
|
||||
if get_tensor_model_parallel_world_size() > 1:
|
||||
rank = get_tensor_model_parallel_rank()
|
||||
attn_bias = attn_bias[:, rank * self.n_heads:(rank + 1) *
|
||||
self.n_heads, :, :]
|
||||
attn_output = self.attn(q, k, v, attn_bias)
|
||||
output, _ = self.o(attn_output)
|
||||
return output
|
||||
|
||||
@@ -52,7 +52,6 @@ def get_hf_config(
|
||||
trust_remote_code: bool,
|
||||
revision: Optional[str] = None,
|
||||
model_override_args: Optional[dict] = None,
|
||||
inference_args: Optional[dict] = None,
|
||||
**kwargs,
|
||||
):
|
||||
is_gguf = check_gguf_file(model)
|
||||
@@ -84,20 +83,23 @@ def get_hf_config(
|
||||
|
||||
def get_diffusers_config(
|
||||
model: str,
|
||||
inference_args: Optional[dict] = None,
|
||||
fastvideo_args: Optional[dict] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Gets a configuration for the given diffusers model.
|
||||
|
||||
Args:
|
||||
model: The model name or path.
|
||||
inference_args: Optional inference arguments to override in the config.
|
||||
fastvideo_args: Optional inference arguments to override in the config.
|
||||
|
||||
Returns:
|
||||
The loaded configuration.
|
||||
"""
|
||||
config_name = "config.json"
|
||||
if "scheduler" in model:
|
||||
config_name = "scheduler_config.json"
|
||||
# Check if the model path exists
|
||||
if os.path.exists(model):
|
||||
config_file = os.path.join(model, "config.json")
|
||||
config_file = os.path.join(model, config_name)
|
||||
if os.path.exists(config_file):
|
||||
try:
|
||||
# Load the config directly from the file
|
||||
|
||||
@@ -10,10 +10,10 @@ from typing import Any, Generator, Iterable, List, Optional, Tuple, cast
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
from transformers import AutoTokenizer, PretrainedConfig
|
||||
from transformers import AutoImageProcessor, AutoTokenizer, PretrainedConfig
|
||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.hf_transformer_utils import (get_diffusers_config,
|
||||
get_hf_config)
|
||||
@@ -23,6 +23,7 @@ from fastvideo.v1.models.loader.weight_utils import (
|
||||
filter_duplicate_safetensors_files, filter_files_not_needed_for_inference,
|
||||
pt_weights_iterator, safetensors_weights_iterator)
|
||||
from fastvideo.v1.models.registry import ModelRegistry
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -35,14 +36,14 @@ class ComponentLoader(ABC):
|
||||
|
||||
@abstractmethod
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Load the component based on the model path, architecture, and inference args.
|
||||
|
||||
Args:
|
||||
model_path: Path to the component model
|
||||
architecture: Architecture of the component model
|
||||
inference_args: Inference arguments
|
||||
fastvideo_args: Inference arguments
|
||||
|
||||
Returns:
|
||||
The loaded component
|
||||
@@ -71,6 +72,8 @@ class ComponentLoader(ABC):
|
||||
"text_encoder_2": (TextEncoderLoader, "transformers"),
|
||||
"tokenizer": (TokenizerLoader, "transformers"),
|
||||
"tokenizer_2": (TokenizerLoader, "transformers"),
|
||||
"image_processor": (ImageProcessorLoader, "transformers"),
|
||||
"image_encoder": (ImageEncoderLoader, "transformers"),
|
||||
}
|
||||
|
||||
if module_type in module_loaders:
|
||||
@@ -196,24 +199,27 @@ class TextEncoderLoader(ComponentLoader):
|
||||
yield from self._get_weights_iterator(source)
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the text encoders based on the model path, architecture, and inference args."""
|
||||
model_config: PretrainedConfig = get_hf_config(
|
||||
model=model_path,
|
||||
trust_remote_code=inference_args.trust_remote_code,
|
||||
revision=inference_args.revision,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
revision=fastvideo_args.revision,
|
||||
model_override_args=None,
|
||||
inference_args=inference_args,
|
||||
)
|
||||
logger.info("HF Model config: %s", model_config)
|
||||
|
||||
target_device = torch.device(inference_args.device_str)
|
||||
target_device = torch.device(fastvideo_args.device_str)
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(model_path, model_config, target_device)
|
||||
return self.load_model(model_path, model_config, target_device,
|
||||
fastvideo_args.text_encoder_precision)
|
||||
|
||||
def load_model(self, model_path: str, model_config,
|
||||
target_device: torch.device):
|
||||
with set_default_torch_dtype(torch.float16):
|
||||
def load_model(self,
|
||||
model_path: str,
|
||||
model_config,
|
||||
target_device: torch.device,
|
||||
dtype: str = "fp16"):
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
|
||||
with target_device:
|
||||
architectures = getattr(model_config, "architectures", [])
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
|
||||
@@ -240,11 +246,44 @@ class TextEncoderLoader(ComponentLoader):
|
||||
return model.eval()
|
||||
|
||||
|
||||
class ImageEncoderLoader(TextEncoderLoader):
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the text encoders based on the model path, architecture, and inference args."""
|
||||
model_config: PretrainedConfig = get_hf_config(
|
||||
model=model_path,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
revision=fastvideo_args.revision,
|
||||
model_override_args=None,
|
||||
)
|
||||
logger.info("HF Model config: %s", model_config)
|
||||
|
||||
target_device = torch.device(fastvideo_args.device_str)
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(model_path, model_config, target_device,
|
||||
fastvideo_args.image_encoder_precision)
|
||||
|
||||
|
||||
class ImageProcessorLoader(ComponentLoader):
|
||||
"""Loader for image processor."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the image processor based on the model path, architecture, and inference args."""
|
||||
logger.info("Loading image processor from %s", model_path)
|
||||
|
||||
image_processor = AutoImageProcessor.from_pretrained(model_path, )
|
||||
logger.info("Loaded image processor: %s",
|
||||
image_processor.__class__.__name__)
|
||||
return image_processor
|
||||
|
||||
|
||||
class TokenizerLoader(ComponentLoader):
|
||||
"""Loader for tokenizers."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the tokenizer based on the model path, architecture, and inference args."""
|
||||
logger.info("Loading tokenizer from %s", model_path)
|
||||
|
||||
@@ -262,11 +301,9 @@ class VAELoader(ComponentLoader):
|
||||
"""Loader for VAE."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the VAE based on the model path, architecture, and inference args."""
|
||||
# TODO(will): move this to a constants file
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
config = get_diffusers_config(model=model_path)
|
||||
|
||||
class_name = config.pop("_class_name")
|
||||
@@ -275,7 +312,7 @@ class VAELoader(ComponentLoader):
|
||||
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
|
||||
vae = vae_cls(**config).to(inference_args.device)
|
||||
vae = vae_cls(**config).to(fastvideo_args.device)
|
||||
|
||||
# Find all safetensors files
|
||||
safetensors_list = glob.glob(
|
||||
@@ -286,17 +323,9 @@ class VAELoader(ComponentLoader):
|
||||
) == 1, f"Found {len(safetensors_list)} safetensors files in {model_path}"
|
||||
loaded = safetensors_load_file(safetensors_list[0])
|
||||
vae.load_state_dict(loaded)
|
||||
dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
|
||||
dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
|
||||
vae = vae.eval().to(dtype)
|
||||
|
||||
# TODO(will): should we define hunyuan vae config class?
|
||||
vae_kwargs = {
|
||||
"s_ratio": config["spatial_compression_ratio"],
|
||||
"t_ratio": config["temporal_compression_ratio"],
|
||||
}
|
||||
|
||||
vae.kwargs = vae_kwargs
|
||||
|
||||
return vae
|
||||
|
||||
|
||||
@@ -304,7 +333,7 @@ class TransformerLoader(ComponentLoader):
|
||||
"""Loader for transformer."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the transformer based on the model path, architecture, and inference args."""
|
||||
model_config = get_diffusers_config(model=model_path)
|
||||
cls_name = model_config.pop("_class_name")
|
||||
@@ -325,20 +354,25 @@ class TransformerLoader(ComponentLoader):
|
||||
logger.info("Loading model from %s safetensors files in %s",
|
||||
len(safetensors_list), model_path)
|
||||
|
||||
# initialize_sequence_parallel_group(inference_args.sp_size)
|
||||
# initialize_sequence_parallel_group(fastvideo_args.sp_size)
|
||||
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
|
||||
# Load the model using FSDP loader
|
||||
logger.info("Loading model from %s", cls_name)
|
||||
model = load_fsdp_model(model_cls=model_cls,
|
||||
init_params=model_config,
|
||||
weight_dir_list=safetensors_list,
|
||||
device=inference_args.device,
|
||||
cpu_offload=inference_args.use_cpu_offload)
|
||||
device=fastvideo_args.device,
|
||||
cpu_offload=fastvideo_args.use_cpu_offload,
|
||||
default_dtype=default_dtype)
|
||||
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
|
||||
|
||||
model.eval()
|
||||
dtypes = set(param.dtype for param in model.parameters())
|
||||
if len(dtypes) > 1:
|
||||
model = model.to(default_dtype)
|
||||
model = model.eval()
|
||||
return model
|
||||
|
||||
|
||||
@@ -346,23 +380,19 @@ class SchedulerLoader(ComponentLoader):
|
||||
"""Loader for scheduler."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the scheduler based on the model path, architecture, and inference args."""
|
||||
if hasattr(inference_args,
|
||||
'denoise_type') and inference_args.denoise_type == "flow":
|
||||
# TODO(will): add schedulers to register or create a new scheduler registry
|
||||
# TODO(will): default to config file but allow override through
|
||||
# inference args. Currently only uses inference args.
|
||||
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchDiscreteScheduler)
|
||||
scheduler = FlowMatchDiscreteScheduler(
|
||||
shift=inference_args.flow_shift,
|
||||
solver=inference_args.flow_solver,
|
||||
)
|
||||
logger.info("Scheduler loaded: %s", scheduler)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid denoise type: {inference_args.denoise_type}")
|
||||
config = get_diffusers_config(model=model_path)
|
||||
|
||||
class_name = config.pop("_class_name")
|
||||
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
|
||||
config.pop("_diffusers_version")
|
||||
|
||||
scheduler_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
|
||||
scheduler = scheduler_cls(**config)
|
||||
if fastvideo_args.flow_shift is not None:
|
||||
scheduler.set_shift(fastvideo_args.flow_shift)
|
||||
|
||||
return scheduler
|
||||
|
||||
@@ -375,7 +405,7 @@ class GenericComponentLoader(ComponentLoader):
|
||||
self.library = library
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load a generic component based on the model path, architecture, and inference args."""
|
||||
logger.warning("Using generic loader for %s with library %s",
|
||||
model_path, self.library)
|
||||
@@ -385,8 +415,8 @@ class GenericComponentLoader(ComponentLoader):
|
||||
|
||||
model = AutoModel.from_pretrained(
|
||||
model_path,
|
||||
trust_remote_code=inference_args.trust_remote_code,
|
||||
revision=inference_args.revision,
|
||||
trust_remote_code=fastvideo_args.trust_remote_code,
|
||||
revision=fastvideo_args.revision,
|
||||
)
|
||||
logger.info("Loaded generic transformers model: %s",
|
||||
model.__class__.__name__)
|
||||
@@ -413,7 +443,7 @@ class PipelineComponentLoader:
|
||||
@staticmethod
|
||||
def load_module(module_name: str, component_model_path: str,
|
||||
transformers_or_diffusers: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Load a pipeline module.
|
||||
|
||||
@@ -422,7 +452,7 @@ class PipelineComponentLoader:
|
||||
component_model_path: Path to the component model
|
||||
transformers_or_diffusers: Whether the module is from transformers or diffusers
|
||||
architecture: Architecture of the component model
|
||||
inference_args: Inference arguments
|
||||
fastvideo_args: Inference arguments
|
||||
|
||||
Returns:
|
||||
The loaded module
|
||||
@@ -439,4 +469,4 @@ class PipelineComponentLoader:
|
||||
transformers_or_diffusers)
|
||||
|
||||
# Load the module
|
||||
return loader.load(component_model_path, architecture, inference_args)
|
||||
return loader.load(component_model_path, architecture, fastvideo_args)
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/parameter.py
|
||||
|
||||
from fractions import Fraction
|
||||
from typing import Any, Callable, Optional, Tuple, Union
|
||||
from typing import Any, Callable, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch.nn import Parameter
|
||||
@@ -58,21 +58,22 @@ class BasevLLMParameter(Parameter):
|
||||
cond2 = loaded_weight.ndim == 0 and loaded_weight.numel() == 1
|
||||
return (cond1 and cond2)
|
||||
|
||||
def _assert_and_load(self, loaded_weight: torch.Tensor):
|
||||
def _assert_and_load(self, loaded_weight: torch.Tensor) -> None:
|
||||
assert (self.data.shape == loaded_weight.shape
|
||||
or self._is_1d_and_scalar(loaded_weight))
|
||||
self.data.copy_(loaded_weight)
|
||||
|
||||
def load_column_parallel_weight(self, loaded_weight: torch.Tensor):
|
||||
def load_column_parallel_weight(self, loaded_weight: torch.Tensor) -> None:
|
||||
self._assert_and_load(loaded_weight)
|
||||
|
||||
def load_row_parallel_weight(self, loaded_weight: torch.Tensor):
|
||||
def load_row_parallel_weight(self, loaded_weight: torch.Tensor) -> None:
|
||||
self._assert_and_load(loaded_weight)
|
||||
|
||||
def load_merged_column_weight(self, loaded_weight: torch.Tensor, **kwargs):
|
||||
def load_merged_column_weight(self, loaded_weight: torch.Tensor,
|
||||
**kwargs) -> None:
|
||||
self._assert_and_load(loaded_weight)
|
||||
|
||||
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs):
|
||||
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs) -> None:
|
||||
self._assert_and_load(loaded_weight)
|
||||
|
||||
|
||||
@@ -95,7 +96,7 @@ class _ColumnvLLMParameter(BasevLLMParameter):
|
||||
def output_dim(self):
|
||||
return self._output_dim
|
||||
|
||||
def load_column_parallel_weight(self, loaded_weight: torch.Tensor):
|
||||
def load_column_parallel_weight(self, loaded_weight: torch.Tensor) -> None:
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
shard_size = self.data.shape[self.output_dim]
|
||||
loaded_weight = loaded_weight.narrow(self.output_dim,
|
||||
@@ -103,10 +104,13 @@ class _ColumnvLLMParameter(BasevLLMParameter):
|
||||
assert self.data.shape == loaded_weight.shape
|
||||
self.data.copy_(loaded_weight)
|
||||
|
||||
def load_merged_column_weight(self, loaded_weight: torch.Tensor, **kwargs):
|
||||
def load_merged_column_weight(self, loaded_weight: torch.Tensor,
|
||||
**kwargs) -> None:
|
||||
|
||||
shard_offset = kwargs.get("shard_offset")
|
||||
shard_size = kwargs.get("shard_size")
|
||||
if shard_offset is None or shard_size is None:
|
||||
raise ValueError("shard_offset and shard_size must be provided")
|
||||
if isinstance(
|
||||
self,
|
||||
(PackedColumnParameter,
|
||||
@@ -124,13 +128,18 @@ class _ColumnvLLMParameter(BasevLLMParameter):
|
||||
assert param_data.shape == loaded_weight.shape
|
||||
param_data.copy_(loaded_weight)
|
||||
|
||||
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs):
|
||||
def load_qkv_weight(self, loaded_weight: torch.Tensor, **kwargs) -> None:
|
||||
|
||||
shard_offset = kwargs.get("shard_offset")
|
||||
shard_size = kwargs.get("shard_size")
|
||||
shard_id = kwargs.get("shard_id")
|
||||
num_heads = kwargs.get("num_heads")
|
||||
|
||||
assert shard_offset is not None
|
||||
assert shard_size is not None
|
||||
assert shard_id is not None
|
||||
assert num_heads is not None
|
||||
|
||||
if isinstance(
|
||||
self,
|
||||
(PackedColumnParameter,
|
||||
@@ -166,7 +175,7 @@ class RowvLLMParameter(BasevLLMParameter):
|
||||
def input_dim(self):
|
||||
return self._input_dim
|
||||
|
||||
def load_row_parallel_weight(self, loaded_weight: torch.Tensor):
|
||||
def load_row_parallel_weight(self, loaded_weight: torch.Tensor) -> None:
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
shard_size = self.data.shape[self.input_dim]
|
||||
loaded_weight = loaded_weight.narrow(self.input_dim,
|
||||
@@ -233,16 +242,16 @@ class PerTensorScaleParameter(BasevLLMParameter):
|
||||
|
||||
# For row parallel layers, no sharding needed
|
||||
# load weight into parameter as is
|
||||
def load_row_parallel_weight(self, *args, **kwargs):
|
||||
def load_row_parallel_weight(self, *args, **kwargs) -> None:
|
||||
super().load_row_parallel_weight(*args, **kwargs)
|
||||
|
||||
def load_merged_column_weight(self, *args, **kwargs):
|
||||
def load_merged_column_weight(self, *args, **kwargs) -> None:
|
||||
self._load_into_shard_id(*args, **kwargs)
|
||||
|
||||
def load_qkv_weight(self, *args, **kwargs):
|
||||
def load_qkv_weight(self, *args, **kwargs) -> None:
|
||||
self._load_into_shard_id(*args, **kwargs)
|
||||
|
||||
def load_column_parallel_weight(self, *args, **kwargs):
|
||||
def load_column_parallel_weight(self, *args, **kwargs) -> None:
|
||||
super().load_row_parallel_weight(*args, **kwargs)
|
||||
|
||||
def _load_into_shard_id(self, loaded_weight: torch.Tensor,
|
||||
@@ -273,14 +282,10 @@ class PackedColumnParameter(_ColumnvLLMParameter):
|
||||
for more details on the packed properties.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
packed_factor: Union[int, Fraction],
|
||||
packed_dim: int,
|
||||
marlin_tile_size: Optional[int] = None,
|
||||
def __init__(self, packed_factor: Union[int, Fraction], packed_dim: int,
|
||||
**kwargs):
|
||||
self._packed_factor = packed_factor
|
||||
self._packed_dim = packed_dim
|
||||
self._marlin_tile_size = marlin_tile_size
|
||||
super().__init__(**kwargs)
|
||||
|
||||
@property
|
||||
@@ -291,17 +296,12 @@ class PackedColumnParameter(_ColumnvLLMParameter):
|
||||
def packed_factor(self):
|
||||
return self._packed_factor
|
||||
|
||||
@property
|
||||
def marlin_tile_size(self):
|
||||
return self._marlin_tile_size
|
||||
|
||||
def adjust_shard_indexes_for_packing(self, shard_size,
|
||||
shard_offset) -> Tuple[Any, Any]:
|
||||
return _adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size,
|
||||
shard_offset=shard_offset,
|
||||
packed_factor=self.packed_factor,
|
||||
marlin_tile_size=self.marlin_tile_size)
|
||||
packed_factor=self.packed_factor)
|
||||
|
||||
|
||||
class PackedvLLMParameter(ModelWeightParameter):
|
||||
@@ -315,14 +315,10 @@ class PackedvLLMParameter(ModelWeightParameter):
|
||||
by accounting for packing and optionally, marlin tile size.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
packed_factor: Union[int, Fraction],
|
||||
packed_dim: int,
|
||||
marlin_tile_size: Optional[int] = None,
|
||||
def __init__(self, packed_factor: Union[int, Fraction], packed_dim: int,
|
||||
**kwargs):
|
||||
self._packed_factor = packed_factor
|
||||
self._packed_dim = packed_dim
|
||||
self._marlin_tile_size = marlin_tile_size
|
||||
super().__init__(**kwargs)
|
||||
|
||||
@property
|
||||
@@ -333,16 +329,11 @@ class PackedvLLMParameter(ModelWeightParameter):
|
||||
def packed_factor(self):
|
||||
return self._packed_factor
|
||||
|
||||
@property
|
||||
def marlin_tile_size(self):
|
||||
return self._marlin_tile_size
|
||||
|
||||
def adjust_shard_indexes_for_packing(self, shard_size, shard_offset):
|
||||
return _adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size,
|
||||
shard_offset=shard_offset,
|
||||
packed_factor=self.packed_factor,
|
||||
marlin_tile_size=self.marlin_tile_size)
|
||||
packed_factor=self.packed_factor)
|
||||
|
||||
|
||||
class BlockQuantScaleParameter(_ColumnvLLMParameter, RowvLLMParameter):
|
||||
@@ -412,18 +403,8 @@ def permute_param_layout_(param: BasevLLMParameter, input_dim: int,
|
||||
return param
|
||||
|
||||
|
||||
def _adjust_shard_indexes_for_marlin(shard_size, shard_offset,
|
||||
marlin_tile_size) -> Tuple[Any, Any]:
|
||||
return shard_size * marlin_tile_size, shard_offset * marlin_tile_size
|
||||
|
||||
|
||||
def _adjust_shard_indexes_for_packing(shard_size, shard_offset, packed_factor,
|
||||
marlin_tile_size) -> Tuple[Any, Any]:
|
||||
def _adjust_shard_indexes_for_packing(shard_size, shard_offset,
|
||||
packed_factor) -> Tuple[Any, Any]:
|
||||
shard_size = shard_size // packed_factor
|
||||
shard_offset = shard_offset // packed_factor
|
||||
if marlin_tile_size is not None:
|
||||
return _adjust_shard_indexes_for_marlin(
|
||||
shard_size=shard_size,
|
||||
shard_offset=shard_offset,
|
||||
marlin_tile_size=marlin_tile_size)
|
||||
return shard_size, shard_offset
|
||||
|
||||
@@ -38,6 +38,7 @@ _TEXT_ENCODER_MODELS = {
|
||||
|
||||
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
|
||||
# "HunyuanVideoTransformer3DModel": ("image_encoder", "hunyuanvideo", "HunyuanVideoImageEncoder"),
|
||||
"CLIPVisionModelWithProjection": ("encoders", "clip", "CLIPVisionModel"),
|
||||
}
|
||||
|
||||
_VAE_MODELS = {
|
||||
@@ -46,12 +47,21 @@ _VAE_MODELS = {
|
||||
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
|
||||
}
|
||||
|
||||
_SCHEDULERS = {
|
||||
"FlowMatchEulerDiscreteScheduler":
|
||||
("schedulers", "scheduling_flow_match_euler_discrete",
|
||||
"FlowMatchDiscreteScheduler"),
|
||||
"UniPCMultistepScheduler":
|
||||
("schedulers", "scheduling_unipc_multistep", "UniPCMultistepScheduler"),
|
||||
}
|
||||
|
||||
_FAST_VIDEO_MODELS = {
|
||||
**_TEXT_TO_VIDEO_DIT_MODELS,
|
||||
**_IMAGE_TO_VIDEO_DIT_MODELS,
|
||||
**_TEXT_ENCODER_MODELS,
|
||||
**_IMAGE_ENCODER_MODELS,
|
||||
**_VAE_MODELS,
|
||||
**_SCHEDULERS,
|
||||
}
|
||||
|
||||
_SUBPROCESS_COMMAND = [
|
||||
|
||||
@@ -0,0 +1,46 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from diffusers.utils import BaseOutput
|
||||
|
||||
|
||||
class BaseScheduler(ABC):
|
||||
timesteps: torch.Tensor
|
||||
order: int
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
# Check if subclass has defined all required properties
|
||||
required_attributes = ['timesteps', 'order']
|
||||
|
||||
for attr in required_attributes:
|
||||
if not hasattr(self, attr):
|
||||
raise AttributeError(
|
||||
f"Subclasses of BaseScheduler must define '{attr}' property"
|
||||
)
|
||||
|
||||
@abstractmethod
|
||||
def set_shift(self, shift: float) -> None:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def set_timesteps(self, *args, **kwargs) -> None:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def scale_model_input(self,
|
||||
sample: torch.Tensor,
|
||||
timestep: Optional[int] = None) -> torch.Tensor:
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.Tensor,
|
||||
timestep: Union[int, torch.Tensor],
|
||||
sample: torch.Tensor,
|
||||
return_dict: bool = True,
|
||||
) -> Union[BaseOutput, Tuple]:
|
||||
pass
|
||||
@@ -27,6 +27,8 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
|
||||
from fastvideo.v1.models.schedulers.base import BaseScheduler
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
@@ -44,7 +46,7 @@ class FlowMatchDiscreteSchedulerOutput(BaseOutput):
|
||||
prev_sample: torch.FloatTensor
|
||||
|
||||
|
||||
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin, BaseScheduler):
|
||||
"""
|
||||
Euler scheduler.
|
||||
|
||||
@@ -74,6 +76,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
reverse: bool = True,
|
||||
solver: str = "euler",
|
||||
n_tokens: Optional[int] = None,
|
||||
**kwargs,
|
||||
):
|
||||
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
|
||||
|
||||
@@ -94,6 +97,8 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
|
||||
)
|
||||
|
||||
BaseScheduler.__init__(self)
|
||||
|
||||
@property
|
||||
def step_index(self):
|
||||
"""
|
||||
@@ -170,6 +175,9 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
|
||||
return idx
|
||||
|
||||
def set_shift(self, shift: float) -> None:
|
||||
self.config.shift = shift
|
||||
|
||||
def _init_step_index(self, timestep) -> None:
|
||||
if self.begin_index is None:
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
@@ -192,6 +200,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
sample: torch.FloatTensor,
|
||||
return_dict: bool = True,
|
||||
**kwargs,
|
||||
) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,202 +0,0 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from transformers.utils import ModelOutput
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def use_default(value, default) -> Any:
|
||||
return value if value is not None else default
|
||||
|
||||
|
||||
@dataclass
|
||||
class TextEncoderModelOutput(ModelOutput):
|
||||
"""
|
||||
Base class for model's outputs that also contains a pooling of the last hidden states.
|
||||
|
||||
Args:
|
||||
hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
|
||||
Sequence of hidden-states at the output of the last layer of the model.
|
||||
|
||||
attention_mask (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
|
||||
Mask to avoid performing attention on padding token indices. Mask values selected in ``[0, 1]``:
|
||||
hidden_states_list (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed):
|
||||
Tuple of `torch.FloatTensor` (one for the output of the embeddings, if the model has an embedding layer, +
|
||||
one for the output of each layer) of shape `(batch_size, sequence_length, hidden_size)`.
|
||||
Hidden-states of the model at the output of each layer plus the optional initial embedding outputs.
|
||||
text_outputs (`list`, *optional*, returned when `return_texts=True` is passed):
|
||||
List of decoded texts.
|
||||
"""
|
||||
|
||||
hidden_state: torch.FloatTensor = None
|
||||
attention_mask: Optional[torch.LongTensor] = None
|
||||
text_outputs: Optional[list] = None
|
||||
|
||||
|
||||
class TextEncoder(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
text_encoder,
|
||||
tokenizer,
|
||||
max_length: int,
|
||||
text_encoder_precision: Optional[str] = None,
|
||||
text_encoder_path: Optional[str] = None,
|
||||
output_key: Optional[str] = None,
|
||||
use_attention_mask: bool = True,
|
||||
prompt_template: Optional[dict] = None,
|
||||
prompt_template_video: Optional[dict] = None,
|
||||
hidden_state_skip_layer: Optional[int] = None,
|
||||
apply_final_norm: bool = False,
|
||||
device=None,
|
||||
):
|
||||
super().__init__()
|
||||
# TODO(will): check if there's a cleaner way to do this
|
||||
self.text_encoder_type = text_encoder.config.architectures[0]
|
||||
self.max_length = max_length
|
||||
self.precision = text_encoder_precision
|
||||
self.model_path = text_encoder_path
|
||||
self.use_attention_mask = use_attention_mask
|
||||
if prompt_template_video is not None:
|
||||
assert (use_attention_mask is True
|
||||
), "Attention mask is True required when training videos."
|
||||
self.prompt_template = prompt_template
|
||||
self.prompt_template_video = prompt_template_video
|
||||
self.hidden_state_skip_layer = hidden_state_skip_layer
|
||||
self.apply_final_norm = apply_final_norm
|
||||
|
||||
if "T5" in self.text_encoder_type:
|
||||
self.output_key = output_key or "last_hidden_state"
|
||||
elif "CLIPTextModel" in self.text_encoder_type:
|
||||
self.output_key = output_key or "pooler_output"
|
||||
elif "LlamaModel" in self.text_encoder_type or "glm" in self.text_encoder_type:
|
||||
self.output_key = output_key or "last_hidden_state"
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported text encoder type: {self.text_encoder_type}")
|
||||
|
||||
self.model = text_encoder
|
||||
# self.dtype = self.model.dtype
|
||||
self.device = device
|
||||
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def __repr__(self):
|
||||
return f"{self.text_encoder_type} ({self.precision} - {self.model_path})"
|
||||
|
||||
@staticmethod
|
||||
def apply_text_to_template(text, template, prevent_empty_text=True) -> str:
|
||||
"""
|
||||
Apply text to template.
|
||||
|
||||
Args:
|
||||
text (str): Input text.
|
||||
template (str or list): Template string or list of chat conversation.
|
||||
prevent_empty_text (bool): If True, we will prevent the user text from being empty
|
||||
by adding a space. Defaults to True.
|
||||
"""
|
||||
if isinstance(template, str):
|
||||
# Will send string to tokenizer. Used for llm
|
||||
return template.format(text)
|
||||
else:
|
||||
raise TypeError(f"Unsupported template type: {type(template)}")
|
||||
|
||||
def text2tokens(self, text) -> dict:
|
||||
"""
|
||||
Tokenize the input text.
|
||||
|
||||
Args:
|
||||
text (str or list): Input text.
|
||||
"""
|
||||
if self.prompt_template_video is not None:
|
||||
prompt_template = self.prompt_template_video["template"]
|
||||
|
||||
text = self.apply_text_to_template(text, prompt_template)
|
||||
|
||||
kwargs = dict(
|
||||
truncation=True,
|
||||
max_length=self.max_length,
|
||||
return_tensors="pt",
|
||||
)
|
||||
batch_encoding: dict = self.tokenizer(
|
||||
text,
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
return_attention_mask=True,
|
||||
**kwargs,
|
||||
)
|
||||
return batch_encoding
|
||||
|
||||
def encode(
|
||||
self,
|
||||
batch_encoding,
|
||||
use_attention_mask=None,
|
||||
hidden_state_skip_layer=None,
|
||||
device=None,
|
||||
) -> TextEncoderModelOutput:
|
||||
"""
|
||||
Args:
|
||||
batch_encoding (dict): Batch encoding from tokenizer.
|
||||
use_attention_mask (bool): Whether to use attention mask. If None, use self.use_attention_mask.
|
||||
Defaults to None.
|
||||
output_hidden_states (bool): Whether to output hidden states. If False, return the value of
|
||||
self.output_key. If True, return the entire output. If set self.hidden_state_skip_layer,
|
||||
output_hidden_states will be set True. Defaults to False.
|
||||
hidden_state_skip_layer (int): Number of hidden states to hidden_state_skip_layer. 0 means the last layer.
|
||||
If None, self.output_key will be used. Defaults to None.
|
||||
return_texts (bool): Whether to return the decoded texts. Defaults to False.
|
||||
"""
|
||||
device = self.model.device if device is None else device
|
||||
use_attention_mask = use_default(use_attention_mask,
|
||||
self.use_attention_mask)
|
||||
hidden_state_skip_layer = use_default(hidden_state_skip_layer,
|
||||
self.hidden_state_skip_layer)
|
||||
|
||||
# note: clip will need attention mask
|
||||
# TODO(will): unify interface with dit
|
||||
# TODO (peiyuan): why clip need attention mask?
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs = self.model(
|
||||
input_ids=batch_encoding["input_ids"].to(device),
|
||||
output_hidden_states=hidden_state_skip_layer is not None,
|
||||
)
|
||||
if hidden_state_skip_layer is not None:
|
||||
last_hidden_state = outputs.hidden_states[-(
|
||||
hidden_state_skip_layer + 1)]
|
||||
# Real last hidden state already has layer norm applied. So here we only apply it
|
||||
# for intermediate layers.
|
||||
if hidden_state_skip_layer > 0 and self.apply_final_norm:
|
||||
last_hidden_state = self.model.final_layer_norm(
|
||||
last_hidden_state)
|
||||
else:
|
||||
last_hidden_state = outputs[self.output_key]
|
||||
|
||||
# Remove hidden states of instruction tokens, only keep prompt tokens.
|
||||
if self.prompt_template_video is not None:
|
||||
|
||||
crop_start = self.prompt_template_video.get("crop_start", -1)
|
||||
|
||||
last_hidden_state = last_hidden_state[:, crop_start:]
|
||||
|
||||
return TextEncoderModelOutput(last_hidden_state)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
text,
|
||||
use_attention_mask=None,
|
||||
output_hidden_states=False,
|
||||
hidden_state_skip_layer=None,
|
||||
return_texts=False,
|
||||
):
|
||||
batch_encoding = self.text2tokens(text)
|
||||
return self.encode(
|
||||
batch_encoding,
|
||||
use_attention_mask=use_attention_mask,
|
||||
hidden_state_skip_layer=hidden_state_skip_layer,
|
||||
)
|
||||
@@ -2,11 +2,12 @@
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from math import prod
|
||||
from typing import Iterator, Optional, Tuple
|
||||
from typing import Iterator, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size)
|
||||
@@ -20,8 +21,11 @@ class ParallelTiledVAE(ABC):
|
||||
tile_sample_stride_width: int
|
||||
tile_sample_stride_num_frames: int
|
||||
use_tiling: bool
|
||||
use_temporal_tiling: bool
|
||||
use_parallel_tiling: bool
|
||||
temporal_compression_ratio: int
|
||||
spatial_compression_ratio: int
|
||||
scaling_factor: Union[float, torch.tensor]
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
# Check if subclass has defined all required properties
|
||||
@@ -30,7 +34,8 @@ class ParallelTiledVAE(ABC):
|
||||
'tile_sample_min_num_frames', 'tile_sample_stride_height',
|
||||
'tile_sample_stride_width', 'tile_sample_stride_num_frames',
|
||||
'spatial_compression_ratio', 'temporal_compression_ratio',
|
||||
'use_tiling'
|
||||
'use_tiling', 'use_temporal_tiling', 'use_parallel_tiling',
|
||||
'scaling_factor'
|
||||
]
|
||||
|
||||
for attr in required_attributes:
|
||||
@@ -52,13 +57,13 @@ class ParallelTiledVAE(ABC):
|
||||
latent_num_frames = (num_frames -
|
||||
1) // self.temporal_compression_ratio + 1
|
||||
|
||||
if self.use_tiling and num_frames > self.tile_sample_min_num_frames:
|
||||
if self.use_tiling and self.use_temporal_tiling and num_frames > self.tile_sample_min_num_frames:
|
||||
latents = self.tiled_encode(x)[:, :, :latent_num_frames]
|
||||
elif self.use_tiling and (width > self.tile_sample_min_width
|
||||
or height > self.tile_sample_min_height):
|
||||
latents = self.spatial_tiled_encode(x)
|
||||
latents = self.spatial_tiled_encode(x)[:, :, :latent_num_frames]
|
||||
else:
|
||||
latents = self._encode(x)
|
||||
latents = self._encode(x)[:, :, :latent_num_frames]
|
||||
return DiagonalGaussianDistribution(latents)
|
||||
|
||||
def decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
@@ -69,16 +74,17 @@ class ParallelTiledVAE(ABC):
|
||||
num_sample_frames = (num_frames -
|
||||
1) * self.temporal_compression_ratio + 1
|
||||
|
||||
if self.use_tiling and get_sequence_model_parallel_world_size() > 1:
|
||||
if self.use_tiling and self.use_parallel_tiling and get_sequence_model_parallel_world_size(
|
||||
) > 1:
|
||||
return self.parallel_tiled_decode(z)[:, :, :num_sample_frames]
|
||||
if self.use_tiling and num_frames > tile_latent_min_num_frames:
|
||||
if self.use_tiling and self.use_temporal_tiling and num_frames > tile_latent_min_num_frames:
|
||||
return self.tiled_decode(z)[:, :, :num_sample_frames]
|
||||
|
||||
if self.use_tiling and (width > tile_latent_min_width
|
||||
or height > tile_latent_min_height):
|
||||
return self.spatial_tiled_decode(z)
|
||||
return self.spatial_tiled_decode(z)[:, :, :num_sample_frames]
|
||||
|
||||
return self._decode(z)
|
||||
return self._decode(z)[:, :, :num_sample_frames]
|
||||
|
||||
def blend_v(self, a: torch.Tensor, b: torch.Tensor,
|
||||
blend_extent: int) -> torch.Tensor:
|
||||
@@ -462,7 +468,7 @@ class DiagonalGaussianDistribution:
|
||||
def sample(self,
|
||||
generator: Optional[torch.Generator] = None) -> torch.Tensor:
|
||||
# make sure sample is on the same device as the parameters and has same dtype
|
||||
sample = torch.randn(
|
||||
sample = randn_tensor(
|
||||
self.mean.shape,
|
||||
generator=generator,
|
||||
device=self.parameters.device,
|
||||
|
||||
@@ -847,6 +847,9 @@ class AutoencoderKLHunyuanVideo(nn.Module, ParallelTiledVAE):
|
||||
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
|
||||
# intermediate tiles together, the memory requirement can be lowered.
|
||||
self.use_tiling = True
|
||||
self.use_temporal_tiling = True
|
||||
self.use_parallel_tiling = True
|
||||
self.scaling_factor = scaling_factor
|
||||
|
||||
# The minimal tile height and width for spatial tiling to be used
|
||||
self.tile_sample_min_height = 256
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
@@ -17,14 +19,34 @@ from typing import Optional, Tuple, Union
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
import torch.utils.checkpoint
|
||||
from contextlib import contextmanager
|
||||
import contextvars
|
||||
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.utils import auto_attributes
|
||||
from fastvideo.v1.models.utils import auto_attributes
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE, DiagonalGaussianDistribution
|
||||
|
||||
CACHE_T = 2
|
||||
|
||||
is_first_frame = contextvars.ContextVar("is_first_frame", default=False)
|
||||
feat_cache = contextvars.ContextVar("feat_cache", default=None)
|
||||
feat_idx = contextvars.ContextVar("feat_idx", default=0)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def forward_context(first_frame_arg=False,
|
||||
feat_cache_arg=None,
|
||||
feat_idx_arg=None):
|
||||
is_first_frame_token = is_first_frame.set(first_frame_arg)
|
||||
feat_cache_token = feat_cache.set(feat_cache_arg)
|
||||
feat_idx_token = feat_idx.set(feat_idx_arg)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
is_first_frame.reset(is_first_frame_token)
|
||||
feat_cache.reset(feat_cache_token)
|
||||
feat_idx.reset(feat_idx_token)
|
||||
|
||||
|
||||
class WanCausalConv3d(nn.Conv3d):
|
||||
r"""
|
||||
@@ -58,12 +80,17 @@ class WanCausalConv3d(nn.Conv3d):
|
||||
)
|
||||
self.padding: Tuple[int, int, int]
|
||||
# Set up causal padding
|
||||
self._padding = (self.padding[2], self.padding[2], self.padding[1],
|
||||
self.padding[1], 2 * self.padding[0], 0)
|
||||
self._padding: Tuple[int, ...] = (self.padding[2], self.padding[2],
|
||||
self.padding[1], self.padding[1],
|
||||
2 * self.padding[0], 0)
|
||||
self.padding = (0, 0, 0)
|
||||
|
||||
def forward(self, x):
|
||||
def forward(self, x, cache_x=None):
|
||||
padding = list(self._padding)
|
||||
if cache_x is not None and self._padding[4] > 0:
|
||||
cache_x = cache_x.to(x.device)
|
||||
x = torch.cat([cache_x, x], dim=2)
|
||||
padding[4] -= cache_x.shape[2]
|
||||
x = F.pad(x, padding)
|
||||
return super().forward(x)
|
||||
|
||||
@@ -155,28 +182,82 @@ class WanResample(nn.Module):
|
||||
self.time_conv = WanCausalConv3d(dim,
|
||||
dim, (3, 1, 1),
|
||||
stride=(2, 1, 1),
|
||||
padding=(1, 0, 0))
|
||||
padding=(0, 0, 0))
|
||||
|
||||
else:
|
||||
self.resample = nn.Identity()
|
||||
|
||||
def forward(self, x, first_frame=False):
|
||||
def forward(self, x):
|
||||
b, c, t, h, w = x.size()
|
||||
first_frame = is_first_frame.get()
|
||||
if first_frame:
|
||||
assert t == 1
|
||||
if self.mode == "upsample3d" and not first_frame and hasattr(
|
||||
self, "time_conv"):
|
||||
x = self.time_conv(x)
|
||||
x = x.reshape(b, 2, c, t, h, w)
|
||||
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3)
|
||||
x = x.reshape(b, c, t * 2, h, w)
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if self.mode == "upsample3d":
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
if _feat_cache[idx] is None:
|
||||
_feat_cache[idx] = "Rep"
|
||||
_feat_idx += 1
|
||||
else:
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and _feat_cache[
|
||||
idx] is not None and _feat_cache[idx] != "Rep":
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
if cache_x.shape[2] < 2 and _feat_cache[
|
||||
idx] is not None and _feat_cache[idx] == "Rep":
|
||||
cache_x = torch.cat([
|
||||
torch.zeros_like(cache_x).to(cache_x.device),
|
||||
cache_x
|
||||
],
|
||||
dim=2)
|
||||
if _feat_cache[idx] == "Rep":
|
||||
x = self.time_conv(x)
|
||||
else:
|
||||
x = self.time_conv(x, _feat_cache[idx])
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
|
||||
x = x.reshape(b, 2, c, t, h, w)
|
||||
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]),
|
||||
3)
|
||||
x = x.reshape(b, c, t * 2, h, w)
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
elif not first_frame and hasattr(self, "time_conv"):
|
||||
x = self.time_conv(x)
|
||||
x = x.reshape(b, 2, c, t, h, w)
|
||||
x = torch.stack((x[:, 0, :, :, :, :], x[:, 1, :, :, :, :]), 3)
|
||||
x = x.reshape(b, c, t * 2, h, w)
|
||||
t = x.shape[2]
|
||||
x = x.permute(0, 2, 1, 3, 4).reshape(b * t, c, h, w)
|
||||
x = self.resample(x)
|
||||
x = x.view(b, t, x.size(1), x.size(2), x.size(3)).permute(0, 2, 1, 3, 4)
|
||||
if self.mode == "downsample3d" and not first_frame and hasattr(
|
||||
self, "time_conv"):
|
||||
x = self.time_conv(x)
|
||||
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if self.mode == "downsample3d":
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
if _feat_cache[idx] is None:
|
||||
_feat_cache[idx] = x.clone()
|
||||
_feat_idx += 1
|
||||
else:
|
||||
cache_x = x[:, :, -1:, :, :].clone()
|
||||
x = self.time_conv(
|
||||
torch.cat([_feat_cache[idx][:, :, -1:, :, :], x], 2))
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
elif not first_frame and hasattr(self, "time_conv"):
|
||||
x = self.time_conv(x)
|
||||
return x
|
||||
|
||||
|
||||
@@ -220,7 +301,25 @@ class WanResidualBlock(nn.Module):
|
||||
x = self.norm1(x)
|
||||
x = self.nonlinearity(x)
|
||||
|
||||
x = self.conv1(x)
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
|
||||
x = self.conv1(x, _feat_cache[idx])
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
else:
|
||||
x = self.conv1(x)
|
||||
|
||||
# Second normalization and activation
|
||||
x = self.norm2(x)
|
||||
@@ -229,7 +328,25 @@ class WanResidualBlock(nn.Module):
|
||||
# Dropout
|
||||
x = self.dropout(x)
|
||||
|
||||
x = self.conv2(x)
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
|
||||
x = self.conv2(x, _feat_cache[idx])
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
else:
|
||||
x = self.conv2(x)
|
||||
|
||||
# Add residual connection
|
||||
return x + h
|
||||
@@ -398,15 +515,30 @@ class WanEncoder3d(nn.Module):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, x, first_frame=False):
|
||||
x = self.conv_in(x)
|
||||
def forward(self, x):
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = self.conv_in(x, _feat_cache[idx])
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
else:
|
||||
x = self.conv_in(x)
|
||||
|
||||
## downsamples
|
||||
for layer in self.down_blocks:
|
||||
if isinstance(layer, WanResample):
|
||||
x = layer(x, first_frame=first_frame)
|
||||
else:
|
||||
x = layer(x)
|
||||
x = layer(x)
|
||||
|
||||
## middle
|
||||
x = self.mid_block(x)
|
||||
@@ -414,7 +546,26 @@ class WanEncoder3d(nn.Module):
|
||||
## head
|
||||
x = self.norm_out(x)
|
||||
x = self.nonlinearity(x)
|
||||
x = self.conv_out(x)
|
||||
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = self.conv_out(x, _feat_cache[idx])
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
else:
|
||||
x = self.conv_out(x)
|
||||
return x
|
||||
|
||||
|
||||
@@ -463,7 +614,7 @@ class WanUpBlock(nn.Module):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, x, first_frame=False):
|
||||
def forward(self, x):
|
||||
"""
|
||||
Forward pass through the upsampling block.
|
||||
|
||||
@@ -479,7 +630,7 @@ class WanUpBlock(nn.Module):
|
||||
x = resnet(x)
|
||||
|
||||
if self.upsamplers is not None:
|
||||
x = self.upsamplers[0](x, first_frame=first_frame)
|
||||
x = self.upsamplers[0](x)
|
||||
return x
|
||||
|
||||
|
||||
@@ -567,21 +718,57 @@ class WanDecoder3d(nn.Module):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(self, x, first_frame=False):
|
||||
def forward(self, x):
|
||||
## conv1
|
||||
x = self.conv_in(x)
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = self.conv_in(x, _feat_cache[idx])
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
else:
|
||||
x = self.conv_in(x)
|
||||
|
||||
## middle
|
||||
x = self.mid_block(x)
|
||||
|
||||
## upsamples
|
||||
for up_block in self.up_blocks:
|
||||
x = up_block(x, first_frame=first_frame)
|
||||
x = up_block(x)
|
||||
|
||||
## head
|
||||
x = self.norm_out(x)
|
||||
x = self.nonlinearity(x)
|
||||
x = self.conv_out(x)
|
||||
_feat_cache = feat_cache.get()
|
||||
_feat_idx = feat_idx.get()
|
||||
if _feat_cache is not None:
|
||||
idx = _feat_idx
|
||||
cache_x = x[:, :, -CACHE_T:, :, :].clone()
|
||||
if cache_x.shape[2] < 2 and _feat_cache[idx] is not None:
|
||||
# cache last frame of last two chunk
|
||||
cache_x = torch.cat([
|
||||
_feat_cache[idx][:, :, -1, :, :].unsqueeze(2).to(
|
||||
cache_x.device), cache_x
|
||||
],
|
||||
dim=2)
|
||||
x = self.conv_out(x, _feat_cache[idx])
|
||||
_feat_cache[idx] = cache_x
|
||||
_feat_idx += 1
|
||||
feat_cache.set(_feat_cache)
|
||||
feat_idx.set(_feat_idx)
|
||||
else:
|
||||
x = self.conv_out(x)
|
||||
return x
|
||||
|
||||
|
||||
@@ -647,6 +834,10 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
self.temperal_upsample = list(temperal_downsample)[::-1]
|
||||
self.latents_mean = list(latents_mean)
|
||||
self.latents_std = list(latents_std)
|
||||
self.scaling_factor = 1.0 / torch.tensor(self.config.latents_std).view(
|
||||
1, self.config.z_dim, 1, 1, 1)
|
||||
self.shift_factor = torch.tensor(self.config.latents_mean).view(
|
||||
1, self.config.z_dim, 1, 1, 1)
|
||||
|
||||
if load_encoder:
|
||||
self.encoder = WanEncoder3d(base_dim, z_dim * 2, dim_mult,
|
||||
@@ -661,6 +852,8 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
self.temperal_upsample, dropout)
|
||||
|
||||
self.use_tiling = True
|
||||
self.use_temporal_tiling = False
|
||||
self.use_parallel_tiling = False
|
||||
self.spatial_compression_ratio = 8
|
||||
self.temporal_compression_ratio = 4
|
||||
|
||||
@@ -673,10 +866,63 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
self.tile_sample_stride_height = 192
|
||||
self.tile_sample_stride_width = 192
|
||||
self.tile_sample_stride_num_frames = 12
|
||||
|
||||
# Whether to use the feature cache algorithm used by diffusers and Wan2.1
|
||||
self.use_feature_cache = True # default to True for best performance
|
||||
ParallelTiledVAE.__init__(self)
|
||||
|
||||
def clear_cache(self) -> None:
|
||||
|
||||
def _count_conv3d(model) -> int:
|
||||
count = 0
|
||||
for m in model.modules():
|
||||
if isinstance(m, WanCausalConv3d):
|
||||
count += 1
|
||||
return count
|
||||
|
||||
self._conv_num = _count_conv3d(self.decoder)
|
||||
self._conv_idx = 0
|
||||
self._feat_map = [None] * self._conv_num
|
||||
# cache encode
|
||||
self._enc_conv_num = _count_conv3d(self.encoder)
|
||||
self._enc_conv_idx = 0
|
||||
self._enc_feat_map = [None] * self._enc_conv_num
|
||||
|
||||
def encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
if self.use_feature_cache:
|
||||
self.clear_cache()
|
||||
with forward_context(feat_cache_arg=self._enc_feat_map,
|
||||
feat_idx_arg=self._enc_conv_idx):
|
||||
t = x.shape[2]
|
||||
iter_ = 1 + (t - 1) // 4
|
||||
for i in range(iter_):
|
||||
feat_idx.set(0)
|
||||
if i == 0:
|
||||
out = self.encoder(x[:, :, :1, :, :])
|
||||
else:
|
||||
out_ = self.encoder(x[:, :,
|
||||
1 + 4 * (i - 1):1 + 4 * i, :, :])
|
||||
out = torch.cat([out, out_], 2)
|
||||
enc = self.quant_conv(out)
|
||||
mu, logvar = enc[:, :self.z_dim, :, :, :], enc[:,
|
||||
self.z_dim:, :, :, :]
|
||||
enc = torch.cat([mu, logvar], dim=1)
|
||||
enc = DiagonalGaussianDistribution(enc)
|
||||
self.clear_cache()
|
||||
else:
|
||||
for block in self.encoder.down_blocks:
|
||||
if isinstance(block,
|
||||
WanResample) and block.mode == "downsample3d":
|
||||
_padding = list(block.time_conv._padding)
|
||||
_padding[4] = 2
|
||||
block.time_conv._padding = tuple(_padding)
|
||||
enc = ParallelTiledVAE.encode(self, x)
|
||||
|
||||
return enc
|
||||
|
||||
def _encode(self, x: torch.Tensor, first_frame=False) -> torch.Tensor:
|
||||
out = self.encoder(x, first_frame=first_frame)
|
||||
with forward_context(first_frame_arg=first_frame):
|
||||
out = self.encoder(x)
|
||||
enc = self.quant_conv(out)
|
||||
mu, logvar = enc[:, :self.z_dim, :, :, :], enc[:, self.z_dim:, :, :, :]
|
||||
enc = torch.cat([mu, logvar], dim=1)
|
||||
@@ -691,14 +937,41 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
enc = torch.cat([first_frame, enc], dim=2)
|
||||
return enc
|
||||
|
||||
def spatial_tiled_encode(self, x: torch.Tensor) -> torch.Tensor:
|
||||
first_frame = x[:, :, 0, :, :].unsqueeze(2)
|
||||
first_frame = self._encode(first_frame, first_frame=True)
|
||||
|
||||
enc = ParallelTiledVAE.spatial_tiled_encode(self, x)
|
||||
enc = enc[:, :, 1:]
|
||||
enc = torch.cat([first_frame, enc], dim=2)
|
||||
return enc
|
||||
|
||||
def decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
if self.use_feature_cache:
|
||||
self.clear_cache()
|
||||
iter_ = z.shape[2]
|
||||
x = self.post_quant_conv(z)
|
||||
with forward_context(feat_cache_arg=self._feat_map,
|
||||
feat_idx_arg=self._conv_idx):
|
||||
for i in range(iter_):
|
||||
feat_idx.set(0)
|
||||
if i == 0:
|
||||
out = self.decoder(x[:, :, i:i + 1, :, :])
|
||||
else:
|
||||
out_ = self.decoder(x[:, :, i:i + 1, :, :])
|
||||
out = torch.cat([out, out_], 2)
|
||||
|
||||
out = torch.clamp(out, min=-1.0, max=1.0)
|
||||
self.clear_cache()
|
||||
else:
|
||||
out = ParallelTiledVAE.decode(self, z)
|
||||
|
||||
return out
|
||||
|
||||
def _decode(self, z: torch.Tensor, first_frame=False) -> torch.Tensor:
|
||||
latents_mean = (torch.tensor(self.latents_mean).view(
|
||||
1, self.z_dim, 1, 1, 1).to(z.device, z.dtype))
|
||||
latents_std = 1.0 / torch.tensor(self.latents_std).view(
|
||||
1, self.z_dim, 1, 1, 1).to(z.device, z.dtype)
|
||||
z = z / latents_std + latents_mean
|
||||
x = self.post_quant_conv(z)
|
||||
out = self.decoder(x, first_frame=first_frame)
|
||||
with forward_context(first_frame_arg=first_frame):
|
||||
out = self.decoder(x)
|
||||
|
||||
out = torch.clamp(out, min=-1.0, max=1.0)
|
||||
|
||||
@@ -711,6 +984,12 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
dec = dec[:, :, start_frame_idx:]
|
||||
return dec
|
||||
|
||||
def spatial_tiled_decode(self, z: torch.Tensor) -> torch.Tensor:
|
||||
dec = ParallelTiledVAE.spatial_tiled_decode(self, z)
|
||||
start_frame_idx = self.temporal_compression_ratio - 1
|
||||
dec = dec[:, :, start_frame_idx:]
|
||||
return dec
|
||||
|
||||
def parallel_tiled_decode(self, z: torch.FloatTensor) -> torch.FloatTensor:
|
||||
self.blend_num_frames *= 2
|
||||
dec = ParallelTiledVAE.parallel_tiled_decode(self, z)
|
||||
|
||||
@@ -0,0 +1,220 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
from typing import Callable, List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
import PIL.ImageOps
|
||||
import requests
|
||||
import torch
|
||||
from packaging import version
|
||||
|
||||
if version.parse(version.parse(
|
||||
PIL.__version__).base_version) >= version.parse("9.1.0"):
|
||||
PIL_INTERPOLATION = {
|
||||
"linear": PIL.Image.Resampling.BILINEAR,
|
||||
"bilinear": PIL.Image.Resampling.BILINEAR,
|
||||
"bicubic": PIL.Image.Resampling.BICUBIC,
|
||||
"lanczos": PIL.Image.Resampling.LANCZOS,
|
||||
"nearest": PIL.Image.Resampling.NEAREST,
|
||||
}
|
||||
else:
|
||||
PIL_INTERPOLATION = {
|
||||
"linear": PIL.Image.LINEAR,
|
||||
"bilinear": PIL.Image.BILINEAR,
|
||||
"bicubic": PIL.Image.BICUBIC,
|
||||
"lanczos": PIL.Image.LANCZOS,
|
||||
"nearest": PIL.Image.NEAREST,
|
||||
}
|
||||
|
||||
|
||||
def pil_to_numpy(
|
||||
images: Union[List[PIL.Image.Image], PIL.Image.Image]) -> np.ndarray:
|
||||
r"""
|
||||
Convert a PIL image or a list of PIL images to NumPy arrays.
|
||||
|
||||
Args:
|
||||
images (`PIL.Image.Image` or `List[PIL.Image.Image]`):
|
||||
The PIL image or list of images to convert to NumPy format.
|
||||
|
||||
Returns:
|
||||
`np.ndarray`:
|
||||
A NumPy array representation of the images.
|
||||
"""
|
||||
if not isinstance(images, list):
|
||||
images = [images]
|
||||
images = [np.array(image).astype(np.float32) / 255.0 for image in images]
|
||||
images_arr: np.ndarray = np.stack(images, axis=0)
|
||||
|
||||
return images_arr
|
||||
|
||||
|
||||
def numpy_to_pt(images: np.ndarray) -> torch.Tensor:
|
||||
r"""
|
||||
Convert a NumPy image to a PyTorch tensor.
|
||||
|
||||
Args:
|
||||
images (`np.ndarray`):
|
||||
The NumPy image array to convert to PyTorch format.
|
||||
|
||||
Returns:
|
||||
`torch.Tensor`:
|
||||
A PyTorch tensor representation of the images.
|
||||
"""
|
||||
if images.ndim == 3:
|
||||
images = images[..., None]
|
||||
|
||||
images = torch.from_numpy(images.transpose(0, 3, 1, 2))
|
||||
return images
|
||||
|
||||
|
||||
def normalize(
|
||||
images: Union[np.ndarray,
|
||||
torch.Tensor]) -> Union[np.ndarray, torch.Tensor]:
|
||||
r"""
|
||||
Normalize an image array to [-1,1].
|
||||
|
||||
Args:
|
||||
images (`np.ndarray` or `torch.Tensor`):
|
||||
The image array to normalize.
|
||||
|
||||
Returns:
|
||||
`np.ndarray` or `torch.Tensor`:
|
||||
The normalized image array.
|
||||
"""
|
||||
return 2.0 * images - 1.0
|
||||
|
||||
|
||||
def load_image(
|
||||
image: Union[str, PIL.Image.Image],
|
||||
convert_method: Optional[Callable[[PIL.Image.Image],
|
||||
PIL.Image.Image]] = None
|
||||
) -> PIL.Image.Image:
|
||||
"""
|
||||
Loads `image` to a PIL Image.
|
||||
|
||||
Args:
|
||||
image (`str` or `PIL.Image.Image`):
|
||||
The image to convert to the PIL Image format.
|
||||
convert_method (Callable[[PIL.Image.Image], PIL.Image.Image], *optional*):
|
||||
A conversion method to apply to the image after loading it. When set to `None` the image will be converted
|
||||
"RGB".
|
||||
|
||||
Returns:
|
||||
`PIL.Image.Image`:
|
||||
A PIL Image.
|
||||
"""
|
||||
if isinstance(image, str):
|
||||
if image.startswith("http://") or image.startswith("https://"):
|
||||
image = PIL.Image.open(requests.get(image, stream=True).raw)
|
||||
elif os.path.isfile(image):
|
||||
image = PIL.Image.open(image)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Incorrect path or URL. URLs must start with `http://` or `https://`, and {image} is not a valid path."
|
||||
)
|
||||
elif isinstance(image, PIL.Image.Image):
|
||||
image = image
|
||||
else:
|
||||
raise ValueError(
|
||||
"Incorrect format used for the image. Should be a URL linking to an image, a local path, or a PIL image."
|
||||
)
|
||||
|
||||
image = PIL.ImageOps.exif_transpose(image)
|
||||
|
||||
if convert_method is not None:
|
||||
image = convert_method(image)
|
||||
else:
|
||||
image = image.convert("RGB")
|
||||
|
||||
return image
|
||||
|
||||
|
||||
def get_default_height_width(
|
||||
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
|
||||
vae_scale_factor: int,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
) -> Tuple[int, int]:
|
||||
r"""
|
||||
Returns the height and width of the image, downscaled to the next integer multiple of `vae_scale_factor`.
|
||||
|
||||
Args:
|
||||
image (`Union[PIL.Image.Image, np.ndarray, torch.Tensor]`):
|
||||
The image input, which can be a PIL image, NumPy array, or PyTorch tensor. If it is a NumPy array, it
|
||||
should have shape `[batch, height, width]` or `[batch, height, width, channels]`. If it is a PyTorch
|
||||
tensor, it should have shape `[batch, channels, height, width]`.
|
||||
height (`Optional[int]`, *optional*, defaults to `None`):
|
||||
The height of the preprocessed image. If `None`, the height of the `image` input will be used.
|
||||
width (`Optional[int]`, *optional*, defaults to `None`):
|
||||
The width of the preprocessed image. If `None`, the width of the `image` input will be used.
|
||||
|
||||
Returns:
|
||||
`Tuple[int, int]`:
|
||||
A tuple containing the height and width, both resized to the nearest integer multiple of
|
||||
`vae_scale_factor`.
|
||||
"""
|
||||
|
||||
if height is None:
|
||||
if isinstance(image, PIL.Image.Image):
|
||||
height = image.height
|
||||
elif isinstance(image, torch.Tensor):
|
||||
height = image.shape[2]
|
||||
else:
|
||||
height = image.shape[1]
|
||||
|
||||
if width is None:
|
||||
if isinstance(image, PIL.Image.Image):
|
||||
width = image.width
|
||||
elif isinstance(image, torch.Tensor):
|
||||
width = image.shape[3]
|
||||
else:
|
||||
width = image.shape[2]
|
||||
|
||||
width, height = (x - x % vae_scale_factor for x in (width, height)
|
||||
) # resize to integer multiple of vae_scale_factor
|
||||
|
||||
return height, width
|
||||
|
||||
|
||||
def resize(
|
||||
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
|
||||
height: int,
|
||||
width: int,
|
||||
resize_mode: str = "default", # "default", "fill", "crop"
|
||||
resample: str = "lanczos",
|
||||
) -> Union[PIL.Image.Image, np.ndarray, torch.Tensor]:
|
||||
"""
|
||||
Resize image.
|
||||
|
||||
Args:
|
||||
image (`PIL.Image.Image`, `np.ndarray` or `torch.Tensor`):
|
||||
The image input, can be a PIL image, numpy array or pytorch tensor.
|
||||
height (`int`):
|
||||
The height to resize to.
|
||||
width (`int`):
|
||||
The width to resize to.
|
||||
resize_mode (`str`, *optional*, defaults to `default`):
|
||||
The resize mode to use, can be one of `default` or `fill`. If `default`, will resize the image to fit
|
||||
within the specified width and height, and it may not maintaining the original aspect ratio. If `fill`,
|
||||
will resize the image to fit within the specified width and height, maintaining the aspect ratio, and
|
||||
then center the image within the dimensions, filling empty with data from image. If `crop`, will resize
|
||||
the image to fit within the specified width and height, maintaining the aspect ratio, and then center
|
||||
the image within the dimensions, cropping the excess. Note that resize_mode `fill` and `crop` are only
|
||||
supported for PIL image input.
|
||||
|
||||
Returns:
|
||||
`PIL.Image.Image`, `np.ndarray` or `torch.Tensor`:
|
||||
The resized image.
|
||||
"""
|
||||
if resize_mode != "default" and not isinstance(image, PIL.Image.Image):
|
||||
raise ValueError(
|
||||
f"Only PIL image input is supported for resize_mode {resize_mode}")
|
||||
assert isinstance(image, PIL.Image.Image)
|
||||
if resize_mode == "default":
|
||||
image = image.resize((width, height),
|
||||
resample=PIL_INTERPOLATION[resample])
|
||||
else:
|
||||
raise ValueError(f"resize_mode {resize_mode} is not supported")
|
||||
return image
|
||||
@@ -40,7 +40,7 @@ from fastvideo.v1.pipelines.stages import (
|
||||
ConditioningStage,
|
||||
# Import other required stages
|
||||
)
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
class YourCustomPipeline(ComposedPipelineBase):
|
||||
@@ -53,7 +53,7 @@ class YourCustomPipeline(ComposedPipelineBase):
|
||||
# Add other required modules
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
# Add and configure pipeline stages
|
||||
self.add_stage(
|
||||
stage_name="input_validation_stage",
|
||||
@@ -61,14 +61,14 @@ class YourCustomPipeline(ComposedPipelineBase):
|
||||
)
|
||||
# Add more stages as needed
|
||||
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# Initialize pipeline-specific components
|
||||
pass
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch, inference_args: InferenceArgs) -> ForwardBatch:
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
# Implement your pipeline's forward pass
|
||||
batch = self.input_validation_stage(batch, inference_args)
|
||||
batch = self.input_validation_stage(batch, fastvideo_args)
|
||||
# Add more stage executions
|
||||
return batch
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ Diffusion pipelines for fastvideo.v1.
|
||||
This package contains diffusion pipelines for generating videos and images.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
@@ -16,7 +16,7 @@ from fastvideo.v1.utils import (maybe_download_model,
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
|
||||
def build_pipeline(fastvideo_args: FastVideoArgs) -> ComposedPipelineBase:
|
||||
"""
|
||||
Only works with valid hf diffusers configs. (model_index.json)
|
||||
We want to build a pipeline based on the inference args mode_path:
|
||||
@@ -25,9 +25,9 @@ def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
|
||||
3. based on the config, determine the pipeline class
|
||||
"""
|
||||
# Get pipeline type
|
||||
model_path = inference_args.model_path
|
||||
model_path = fastvideo_args.model_path
|
||||
model_path = maybe_download_model(model_path)
|
||||
# inference_args.downloaded_model_path = model_path
|
||||
# fastvideo_args.downloaded_model_path = model_path
|
||||
logger.info("Model path: %s", model_path)
|
||||
config = verify_model_config_and_directory(model_path)
|
||||
|
||||
@@ -41,7 +41,7 @@ def build_pipeline(inference_args: InferenceArgs) -> ComposedPipelineBase:
|
||||
pipeline_architecture)
|
||||
|
||||
# instantiate the pipeline
|
||||
pipeline = pipeline_cls(model_path, inference_args, config)
|
||||
pipeline = pipeline_cls(model_path, fastvideo_args, config)
|
||||
logger.info("Pipeline instantiated")
|
||||
|
||||
# pipeline is now initialized and ready to use
|
||||
|
||||
@@ -12,7 +12,7 @@ from typing import Any, Dict, List, Optional, cast
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
@@ -38,7 +38,7 @@ class ComposedPipelineBase(ABC):
|
||||
# TODO(will): args should support both inference args and training args
|
||||
def __init__(self,
|
||||
model_path: str,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
config: Optional[Dict[str, Any]] = None):
|
||||
"""
|
||||
Initialize the pipeline. After __init__, the pipeline should be ready to
|
||||
@@ -61,12 +61,12 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
# Load modules directly in initialization
|
||||
logger.info("Loading pipeline modules...")
|
||||
self.modules = self.load_modules(inference_args)
|
||||
self.modules = self.load_modules(fastvideo_args)
|
||||
|
||||
self.initialize_pipeline(inference_args)
|
||||
self.initialize_pipeline(fastvideo_args)
|
||||
|
||||
logger.info("Creating pipeline stages...")
|
||||
self.create_pipeline_stages(inference_args)
|
||||
self.create_pipeline_stages(fastvideo_args)
|
||||
|
||||
def get_module(self, module_name: str) -> Any:
|
||||
return self.modules[module_name]
|
||||
@@ -77,7 +77,7 @@ class ComposedPipelineBase(ABC):
|
||||
def _load_config(self, model_path: str) -> Dict[str, Any]:
|
||||
model_path = maybe_download_model(self.model_path)
|
||||
self.model_path = model_path
|
||||
# inference_args.downloaded_model_path = model_path
|
||||
# fastvideo_args.downloaded_model_path = model_path
|
||||
logger.info("Model path: %s", model_path)
|
||||
config = verify_model_config_and_directory(model_path)
|
||||
return cast(Dict[str, Any], config)
|
||||
@@ -108,20 +108,20 @@ class ComposedPipelineBase(ABC):
|
||||
return self._stages
|
||||
|
||||
@abstractmethod
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Create the pipeline stages.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def load_modules(self, inference_args: InferenceArgs) -> Dict[str, Any]:
|
||||
def load_modules(self, fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
"""
|
||||
Load the modules from the config.
|
||||
"""
|
||||
@@ -156,7 +156,7 @@ class ComposedPipelineBase(ABC):
|
||||
component_model_path=component_model_path,
|
||||
transformers_or_diffusers=transformers_or_diffusers,
|
||||
architecture=architecture,
|
||||
inference_args=inference_args,
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
logger.info("Loaded module %s from %s", module_name,
|
||||
component_model_path)
|
||||
@@ -185,14 +185,14 @@ class ComposedPipelineBase(ABC):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Generate a video or image using the pipeline.
|
||||
|
||||
Args:
|
||||
batch: The batch to generate from.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
Returns:
|
||||
ForwardBatch: The batch with the generated video or image.
|
||||
"""
|
||||
@@ -201,7 +201,7 @@ class ComposedPipelineBase(ABC):
|
||||
self._stage_name_mapping.keys())
|
||||
logger.info("Batch: %s", batch)
|
||||
for stage in self.stages:
|
||||
batch = stage(batch, inference_args)
|
||||
batch = stage(batch, fastvideo_args)
|
||||
|
||||
# Return the output
|
||||
return batch
|
||||
|
||||
@@ -8,7 +8,7 @@ using the modular pipeline architecture.
|
||||
|
||||
from diffusers.image_processor import VaeImageProcessor
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (CLIPTextEncodingStage,
|
||||
@@ -30,7 +30,7 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
|
||||
"transformer", "scheduler"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, inference_args: InferenceArgs):
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
@@ -67,20 +67,20 @@ class HunyuanVideoPipeline(ComposedPipelineBase):
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def initialize_pipeline(self, inference_args: InferenceArgs):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
vae_scale_factor = 2**(len(self.get_module("vae").block_out_channels) -
|
||||
1)
|
||||
inference_args.vae_scale_factor = vae_scale_factor
|
||||
fastvideo_args.vae_scale_factor = vae_scale_factor
|
||||
|
||||
self.image_processor = VaeImageProcessor(
|
||||
vae_scale_factor=vae_scale_factor)
|
||||
self.add_module("image_processor", self.image_processor)
|
||||
|
||||
num_channels_latents = self.get_module("transformer").in_channels
|
||||
inference_args.num_channels_latents = num_channels_latents
|
||||
fastvideo_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
EntryClass = HunyuanVideoPipeline
|
||||
|
||||
@@ -22,13 +22,17 @@ class ForwardBatch:
|
||||
execution, allowing methods to update specific components without needing
|
||||
to manage numerous individual parameters.
|
||||
"""
|
||||
# TODO(will): double check that args are separate from inference_args
|
||||
# TODO(will): double check that args are separate from fastvideo_args
|
||||
# properly. Also maybe think about providing an abstraction for pipeline
|
||||
# specific arguments.
|
||||
data_type: str
|
||||
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None
|
||||
|
||||
# Image inputs
|
||||
image_path: Optional[str] = None
|
||||
image_embeds: List[torch.Tensor] = field(default_factory=list)
|
||||
|
||||
# Text inputs
|
||||
prompt: Optional[Union[str, List[str]]] = None
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None
|
||||
@@ -36,8 +40,6 @@ class ForwardBatch:
|
||||
# Primary encoder embeddings
|
||||
prompt_embeds: List[torch.Tensor] = field(default_factory=list)
|
||||
negative_prompt_embeds: Optional[List[torch.Tensor]] = None
|
||||
attention_mask: List[torch.Tensor] = field(default_factory=list)
|
||||
negative_attention_mask: List[torch.Tensor] = field(default_factory=list)
|
||||
|
||||
# Additional text-related parameters
|
||||
max_sequence_length: Optional[int] = None
|
||||
@@ -55,6 +57,7 @@ class ForwardBatch:
|
||||
# Latent tensors
|
||||
latents: Optional[torch.Tensor] = None
|
||||
noise_pred: Optional[torch.Tensor] = None
|
||||
image_latent: Optional[torch.Tensor] = None
|
||||
|
||||
# Latent dimensions
|
||||
num_channels_latents: Optional[int] = None
|
||||
@@ -100,3 +103,4 @@ class ForwardBatch:
|
||||
# Set do_classifier_free_guidance based on guidance scale and negative prompt
|
||||
if self.guidance_scale > 1.0:
|
||||
self.do_classifier_free_guidance = True
|
||||
self.negative_prompt_embeds = []
|
||||
|
||||
@@ -7,15 +7,19 @@ complete diffusion pipelines.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.clip_image_encoding import (
|
||||
CLIPImageEncodingStage)
|
||||
from fastvideo.v1.pipelines.stages.clip_text_encoding import (
|
||||
CLIPTextEncodingStage)
|
||||
from fastvideo.v1.pipelines.stages.conditioning import ConditioningStage
|
||||
from fastvideo.v1.pipelines.stages.decoding import DecodingStage
|
||||
from fastvideo.v1.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.v1.pipelines.stages.encoding import EncodingStage
|
||||
from fastvideo.v1.pipelines.stages.input_validation import InputValidationStage
|
||||
from fastvideo.v1.pipelines.stages.latent_preparation import (
|
||||
LatentPreparationStage)
|
||||
from fastvideo.v1.pipelines.stages.llama_encoding import LlamaEncodingStage
|
||||
from fastvideo.v1.pipelines.stages.t5_encoding import T5EncodingStage
|
||||
from fastvideo.v1.pipelines.stages.timestep_preparation import (
|
||||
TimestepPreparationStage)
|
||||
|
||||
@@ -26,7 +30,10 @@ __all__ = [
|
||||
"LatentPreparationStage",
|
||||
"ConditioningStage",
|
||||
"DenoisingStage",
|
||||
"EncodingStage",
|
||||
"DecodingStage",
|
||||
"LlamaEncodingStage",
|
||||
"T5EncodingStage",
|
||||
"CLIPTextEncodingStage",
|
||||
"CLIPImageEncodingStage",
|
||||
]
|
||||
|
||||
@@ -12,7 +12,7 @@ from abc import ABC, abstractmethod
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
@@ -45,7 +45,7 @@ class PipelineStage(ABC):
|
||||
def __call__(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Execute the stage's processing on the batch with optional logging.
|
||||
@@ -53,7 +53,7 @@ class PipelineStage(ABC):
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The updated batch information after this stage's processing.
|
||||
@@ -65,7 +65,7 @@ class PipelineStage(ABC):
|
||||
|
||||
try:
|
||||
# Call the actual implementation
|
||||
result = self._call_implementation(batch, inference_args)
|
||||
result = self._call_implementation(batch, fastvideo_args)
|
||||
|
||||
execution_time = time.time() - start_time
|
||||
self._logger.info("[%s] Execution completed in %s ms",
|
||||
@@ -85,13 +85,13 @@ class PipelineStage(ABC):
|
||||
else:
|
||||
# Just call the implementation directly if logging is disabled
|
||||
# TODO(will): Also handle backward
|
||||
return self.forward(batch, inference_args)
|
||||
return self.forward(batch, fastvideo_args)
|
||||
|
||||
@abstractmethod
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Forward pass of the stage's processing.
|
||||
@@ -101,7 +101,7 @@ class PipelineStage(ABC):
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The updated batch information after this stage's processing.
|
||||
@@ -111,6 +111,6 @@ class PipelineStage(ABC):
|
||||
def backward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
raise NotImplementedError
|
||||
|
||||
@@ -0,0 +1,71 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Image encoding stages for I2V diffusion pipelines.
|
||||
|
||||
This module contains implementations of image encoding stages for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vision_utils import load_image
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CLIPImageEncodingStage(PipelineStage):
|
||||
"""
|
||||
Stage for encoding image prompts into embeddings for diffusion models.
|
||||
|
||||
This stage handles the encoding of image prompts into the embedding space
|
||||
expected by the diffusion model.
|
||||
"""
|
||||
|
||||
def __init__(self, image_encoder, image_processor) -> None:
|
||||
"""
|
||||
Initialize the prompt encoding stage.
|
||||
|
||||
Args:
|
||||
enable_logging: Whether to enable logging for this stage.
|
||||
is_secondary: Whether this is a secondary image encoder.
|
||||
"""
|
||||
super().__init__()
|
||||
self.image_processor = image_processor
|
||||
self.image_encoder = image_encoder
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into image encoder hidden states.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.image_encoder = self.image_encoder.to(batch.device)
|
||||
|
||||
image = load_image(batch.image_path)
|
||||
|
||||
image_inputs = self.image_processor(
|
||||
images=image, return_tensors="pt").to(batch.device)
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
image_embeds = self.image_encoder(**image_inputs)
|
||||
|
||||
batch.image_embeds.append(image_embeds)
|
||||
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.image_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
@@ -5,8 +5,10 @@ Prompt encoding stages for diffusion pipelines.
|
||||
This module contains implementations of prompt encoding stages for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -37,18 +39,20 @@ class CLIPTextEncodingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
batch.prompt,
|
||||
@@ -64,4 +68,25 @@ class CLIPTextEncodingStage(PipelineStage):
|
||||
|
||||
batch.prompt_embeds.append(prompt_embeds)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
negative_text_inputs = self.tokenizer(
|
||||
batch.negative_prompt,
|
||||
truncation=True,
|
||||
# better way to handle this?
|
||||
max_length=77,
|
||||
return_tensors="pt",
|
||||
)
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
negative_outputs = self.text_encoder(
|
||||
input_ids=negative_text_inputs["input_ids"].to(
|
||||
batch.device), )
|
||||
negative_prompt_embeds = negative_outputs["pooler_output"]
|
||||
|
||||
assert batch.negative_prompt_embeds is not None
|
||||
batch.negative_prompt_embeds.append(negative_prompt_embeds)
|
||||
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
|
||||
@@ -5,7 +5,7 @@ Conditioning stage for diffusion pipelines.
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -24,14 +24,14 @@ class ConditioningStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Apply conditioning to the diffusion process.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with applied conditioning.
|
||||
@@ -39,8 +39,7 @@ class ConditioningStage(PipelineStage):
|
||||
if not batch.do_classifier_free_guidance:
|
||||
return batch
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
"Classifier-free guidance is not supported yet")
|
||||
return batch
|
||||
|
||||
logger.info("batch.negative_prompt_embeds: %s",
|
||||
batch.negative_prompt_embeds)
|
||||
|
||||
@@ -5,7 +5,7 @@ Decoding stage for diffusion pipelines.
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -28,14 +28,14 @@ class DecodingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Decode latent representations into pixel space.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with decoded outputs.
|
||||
@@ -46,30 +46,39 @@ class DecodingStage(PipelineStage):
|
||||
raise ValueError("Latents must be provided")
|
||||
|
||||
# Skip decoding if output type is latent
|
||||
if inference_args.output_type == "latent":
|
||||
if fastvideo_args.output_type == "latent":
|
||||
image = latents
|
||||
else:
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
|
||||
vae_autocast_enabled = (vae_dtype != torch.float32
|
||||
) and not inference_args.disable_autocast
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Apply scaling/shifting if needed
|
||||
if (hasattr(self.vae.config, "shift_factor")
|
||||
and self.vae.config.shift_factor):
|
||||
latents = (latents / self.vae.config.scaling_factor +
|
||||
self.vae.config.shift_factor)
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latents = latents / self.vae.scaling_factor.to(
|
||||
latents.device, latents.dtype)
|
||||
else:
|
||||
latents = latents / self.vae.config.scaling_factor
|
||||
latents = latents / self.vae.scaling_factor
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latents += self.vae.shift_factor.to(latents.device,
|
||||
latents.dtype)
|
||||
else:
|
||||
latents += self.vae.shift_factor
|
||||
|
||||
# Decode latents
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if inference_args.vae_tiling:
|
||||
if fastvideo_args.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if inference_args.vae_sp:
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
latents = latents.to(vae_dtype)
|
||||
image = self.vae.decode(latents)
|
||||
|
||||
# Normalize image to [0, 1] range
|
||||
|
||||
@@ -3,6 +3,7 @@
|
||||
Denoising stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
import inspect
|
||||
from typing import Any, Dict, Iterable, Optional
|
||||
|
||||
@@ -16,12 +17,21 @@ from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
|
||||
from fastvideo.v1.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather)
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
st_attn_available = False
|
||||
spec = importlib.util.find_spec("st_attn")
|
||||
if spec is not None:
|
||||
st_attn_available = True
|
||||
|
||||
from fastvideo.v1.attention.backends.sliding_tile_attn import (
|
||||
SlidingTileAttentionBackend)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -41,18 +51,21 @@ class DenoisingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Run the denoising loop.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with denoised latents.
|
||||
"""
|
||||
# If use cpu offload, need to load the model back into gpu again
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.transformer = self.transformer.to(batch.device)
|
||||
# Prepare extra step kwargs for scheduler
|
||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.scheduler.step,
|
||||
@@ -63,9 +76,9 @@ class DenoisingStage(PipelineStage):
|
||||
)
|
||||
|
||||
# Setup precision and autocast settings
|
||||
target_dtype = PRECISION_TO_TYPE[inference_args.precision]
|
||||
target_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not inference_args.disable_autocast
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Handle sequence parallelism if enabled
|
||||
world_size, rank = get_sequence_model_parallel_world_size(
|
||||
@@ -77,6 +90,12 @@ class DenoisingStage(PipelineStage):
|
||||
n=world_size).contiguous()
|
||||
latents = latents[:, :, rank, :, :, :]
|
||||
batch.latents = latents
|
||||
if batch.image_latent is not None:
|
||||
image_latent = rearrange(batch.image_latent,
|
||||
"b t (n s) h w -> b t n s h w",
|
||||
n=world_size).contiguous()
|
||||
image_latent = image_latent[:, :, rank, :, :, :]
|
||||
batch.image_latent = image_latent
|
||||
|
||||
# Get timesteps and calculate warmup steps
|
||||
timesteps = batch.timesteps
|
||||
@@ -98,9 +117,29 @@ class DenoisingStage(PipelineStage):
|
||||
result[t][layer][h] = value
|
||||
return result
|
||||
|
||||
# Prepare image latents and embeddings for I2V generation
|
||||
image_embeds = batch.image_embeds
|
||||
if len(image_embeds) > 0:
|
||||
assert torch.isnan(image_embeds[0]).sum() == 0
|
||||
image_embeds = [
|
||||
image_embed.to(target_dtype) for image_embed in image_embeds
|
||||
]
|
||||
|
||||
image_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
"encoder_hidden_states_image": image_embeds,
|
||||
},
|
||||
)
|
||||
|
||||
# Get latents and embeddings
|
||||
latents = batch.latents
|
||||
prompt_embeds = batch.prompt_embeds
|
||||
assert torch.isnan(prompt_embeds[0]).sum() == 0
|
||||
if batch.do_classifier_free_guidance:
|
||||
neg_prompt_embeds = batch.negative_prompt_embeds
|
||||
assert neg_prompt_embeds is not None
|
||||
assert torch.isnan(neg_prompt_embeds[0]).sum() == 0
|
||||
|
||||
# Run denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
@@ -109,21 +148,24 @@ class DenoisingStage(PipelineStage):
|
||||
if hasattr(self, 'interrupt') and self.interrupt:
|
||||
continue
|
||||
|
||||
# Expand latents for classifier-free guidance
|
||||
latent_model_input = (torch.cat(
|
||||
[latents] *
|
||||
2) if batch.do_classifier_free_guidance else latents)
|
||||
# Expand latents for I2V
|
||||
latent_model_input = latents.to(target_dtype)
|
||||
if batch.image_latent is not None:
|
||||
latent_model_input = torch.cat(
|
||||
[latent_model_input, batch.image_latent],
|
||||
dim=1).to(target_dtype)
|
||||
assert torch.isnan(latent_model_input).sum() == 0
|
||||
latent_model_input = self.scheduler.scale_model_input(
|
||||
latent_model_input, t)
|
||||
|
||||
# Prepare inputs for transformer
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
guidance_expand = (torch.tensor(
|
||||
[inference_args.embedded_cfg_scale] *
|
||||
[fastvideo_args.embedded_cfg_scale] *
|
||||
latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=batch.device,
|
||||
).to(target_dtype) * 1000.0 if inference_args.embedded_cfg_scale
|
||||
).to(target_dtype) * 1000.0 if fastvideo_args.embedded_cfg_scale
|
||||
is not None else None)
|
||||
|
||||
# Predict noise residual
|
||||
@@ -136,43 +178,36 @@ class DenoisingStage(PipelineStage):
|
||||
self.attn_backend = get_attn_backend(
|
||||
head_size=attn_head_size,
|
||||
dtype=torch.float16, # TODO(will): hack
|
||||
distributed=True,
|
||||
supported_attention_backends=[
|
||||
_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA
|
||||
] # hack
|
||||
)
|
||||
|
||||
# TODO(will): clean this up...
|
||||
try:
|
||||
from fastvideo.v1.attention.backends.sliding_tile_attn import (
|
||||
SlidingTileAttentionBackend)
|
||||
except ImportError:
|
||||
SlidingTileAttentionBackend = None
|
||||
|
||||
if SlidingTileAttentionBackend is not None and isinstance(
|
||||
self.attn_backend, SlidingTileAttentionBackend):
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
|
||||
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
|
||||
)
|
||||
if self.attn_metadata_builder_cls is not None:
|
||||
self.attn_metadata_builder = self.attn_metadata_builder_cls(
|
||||
)
|
||||
# TODO(will-refactor): should this be in a new stage?
|
||||
# TODO(will): clean this up
|
||||
attn_metadata = self.attn_metadata_builder.build(
|
||||
current_timestep=i,
|
||||
forward_batch=batch,
|
||||
inference_args=inference_args,
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
assert attn_metadata is not None, "attn_metadata cannot be None"
|
||||
else:
|
||||
attn_metadata = None
|
||||
else:
|
||||
attn_metadata = None
|
||||
|
||||
# TODO(will): finalize the interface. vLLM uses this to
|
||||
# support torch dynamo compilation. They pass in
|
||||
# attn_metadata, vllm_config, and num_tokens. We can pass in
|
||||
# inference_args or training_args, and attn_metadata.
|
||||
# fastvideo_args or training_args, and attn_metadata.
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
# inference_args=inference_args
|
||||
# fastvideo_args=fastvideo_args
|
||||
):
|
||||
# Run transformer
|
||||
noise_pred = self.transformer(
|
||||
@@ -180,29 +215,43 @@ class DenoisingStage(PipelineStage):
|
||||
prompt_embeds,
|
||||
t_expand,
|
||||
guidance=guidance_expand,
|
||||
**image_kwargs,
|
||||
)
|
||||
|
||||
# Apply guidance
|
||||
if batch.do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + batch.guidance_scale * (
|
||||
noise_pred_text - noise_pred_uncond)
|
||||
# Apply guidance
|
||||
if batch.do_classifier_free_guidance:
|
||||
with set_forward_context(
|
||||
current_timestep=i,
|
||||
attn_metadata=attn_metadata,
|
||||
# fastvideo_args=fastvideo_args
|
||||
):
|
||||
# Run transformer
|
||||
noise_pred_uncond = self.transformer(
|
||||
latent_model_input,
|
||||
neg_prompt_embeds,
|
||||
t_expand,
|
||||
guidance=guidance_expand,
|
||||
**image_kwargs,
|
||||
)
|
||||
noise_pred_text = noise_pred
|
||||
noise_pred = noise_pred_uncond + batch.guidance_scale * (
|
||||
noise_pred_text - noise_pred_uncond)
|
||||
|
||||
# Apply guidance rescale if needed
|
||||
if batch.guidance_rescale > 0.0:
|
||||
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
|
||||
noise_pred = self.rescale_noise_cfg(
|
||||
noise_pred,
|
||||
noise_pred_text,
|
||||
guidance_rescale=batch.guidance_rescale,
|
||||
)
|
||||
# Apply guidance rescale if needed
|
||||
if batch.guidance_rescale > 0.0:
|
||||
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
|
||||
noise_pred = self.rescale_noise_cfg(
|
||||
noise_pred,
|
||||
noise_pred_text,
|
||||
guidance_rescale=batch.guidance_rescale,
|
||||
)
|
||||
|
||||
# Compute the previous noisy sample
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
**extra_step_kwargs,
|
||||
return_dict=False)[0]
|
||||
# Compute the previous noisy sample
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
**extra_step_kwargs,
|
||||
return_dict=False)[0]
|
||||
|
||||
# Update progress bar
|
||||
if i == len(timesteps) - 1 or (
|
||||
@@ -218,11 +267,15 @@ class DenoisingStage(PipelineStage):
|
||||
# Update batch with final latents
|
||||
batch.latents = latents
|
||||
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.transformer.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
|
||||
def prepare_extra_func_kwargs(self, func, kwargs) -> Dict[str, Any]:
|
||||
"""
|
||||
Prepare extra kwargs for the scheduler step.
|
||||
Prepare extra kwargs for the scheduler step / denoise step.
|
||||
|
||||
Args:
|
||||
func: The function to prepare kwargs for.
|
||||
|
||||
@@ -0,0 +1,170 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Encoding stage for diffusion pipelines.
|
||||
"""
|
||||
from typing import Optional
|
||||
|
||||
import PIL.Image
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vision_utils import (get_default_height_width,
|
||||
load_image, normalize,
|
||||
numpy_to_pt, pil_to_numpy, resize)
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class EncodingStage(PipelineStage):
|
||||
"""
|
||||
Stage for encoding pixel representations into latent space.
|
||||
|
||||
This stage handles the encoding of pixel representations into the final
|
||||
input format (e.g., latents).
|
||||
"""
|
||||
|
||||
def __init__(self, vae) -> None:
|
||||
self.vae = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode pixel representations into latent space.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded outputs.
|
||||
"""
|
||||
image_path = batch.image_path
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
if image_path is None:
|
||||
raise ValueError("Image Path must be provided")
|
||||
latent_height = batch.height // self.vae.spatial_compression_ratio
|
||||
latent_width = batch.width // self.vae.spatial_compression_ratio
|
||||
|
||||
image = load_image(image_path)
|
||||
image = self.preprocess(
|
||||
image,
|
||||
vae_scale_factor=self.vae.spatial_compression_ratio,
|
||||
height=batch.height,
|
||||
width=batch.width).to(batch.device, dtype=torch.float32)
|
||||
image = image.unsqueeze(2)
|
||||
video_condition = torch.cat([
|
||||
image,
|
||||
image.new_zeros(image.shape[0], image.shape[1],
|
||||
fastvideo_args.num_frames - 1, batch.height,
|
||||
batch.width)
|
||||
],
|
||||
dim=2)
|
||||
video_condition = video_condition.to(device=batch.device,
|
||||
dtype=torch.float32)
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Encode Image
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
if not vae_autocast_enabled:
|
||||
video_condition = video_condition.to(vae_dtype)
|
||||
encoder_output = self.vae.encode(video_condition)
|
||||
|
||||
generator = batch.generator
|
||||
if generator is None:
|
||||
raise ValueError("Generator must be provided")
|
||||
latent_condition = self.retrieve_latents(encoder_output, generator[0])
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
latent_condition -= self.vae.shift_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition -= self.vae.shift_factor
|
||||
|
||||
if isinstance(self.vae.scaling_factor, torch.Tensor):
|
||||
latent_condition = latent_condition * self.vae.scaling_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition = latent_condition * self.vae.scaling_factor
|
||||
|
||||
mask_lat_size = torch.ones(1, 1, fastvideo_args.num_frames,
|
||||
latent_height, latent_width)
|
||||
mask_lat_size[:, :, list(range(1, fastvideo_args.num_frames))] = 0
|
||||
first_frame_mask = mask_lat_size[:, :, 0:1]
|
||||
first_frame_mask = torch.repeat_interleave(
|
||||
first_frame_mask,
|
||||
dim=2,
|
||||
repeats=self.vae.temporal_compression_ratio)
|
||||
mask_lat_size = torch.concat(
|
||||
[first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
|
||||
mask_lat_size = mask_lat_size.view(1, -1,
|
||||
self.vae.temporal_compression_ratio,
|
||||
latent_height, latent_width)
|
||||
mask_lat_size = mask_lat_size.transpose(1, 2)
|
||||
mask_lat_size = mask_lat_size.to(latent_condition.device)
|
||||
|
||||
batch.image_latent = torch.concat([mask_lat_size, latent_condition],
|
||||
dim=1)
|
||||
|
||||
# Offload models if needed
|
||||
if hasattr(self, 'maybe_free_model_hooks'):
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
return batch
|
||||
|
||||
def retrieve_latents(self,
|
||||
encoder_output: torch.Tensor,
|
||||
generator: Optional[torch.Generator] = None,
|
||||
sample_mode: str = "sample"):
|
||||
if sample_mode == "sample":
|
||||
return encoder_output.sample(generator)
|
||||
elif sample_mode == "argmax":
|
||||
return encoder_output.mode()
|
||||
else:
|
||||
raise AttributeError(
|
||||
"Could not access latents of provided encoder_output")
|
||||
|
||||
def preprocess(
|
||||
self,
|
||||
image: PIL.Image.Image,
|
||||
vae_scale_factor: int,
|
||||
height: Optional[int] = None,
|
||||
width: Optional[int] = None,
|
||||
resize_mode: str = "default", # "default", "fill", "crop"
|
||||
) -> torch.Tensor:
|
||||
image = [image]
|
||||
|
||||
height, width = get_default_height_width(image[0], vae_scale_factor,
|
||||
height, width)
|
||||
image = [
|
||||
resize(i, height, width, resize_mode=resize_mode) for i in image
|
||||
]
|
||||
image = pil_to_numpy(image) # to np
|
||||
image = numpy_to_pt(image) # to pt
|
||||
|
||||
do_normalize = True
|
||||
if image.min() < 0:
|
||||
do_normalize = False
|
||||
if do_normalize:
|
||||
image = normalize(image)
|
||||
|
||||
return image
|
||||
@@ -5,7 +5,7 @@ Input validation stage for diffusion pipelines.
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -22,10 +22,10 @@ class InputValidationStage(PipelineStage):
|
||||
"""
|
||||
|
||||
def _generate_seeds(self, batch: ForwardBatch,
|
||||
inference_args: InferenceArgs):
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Generate seeds for the inference"""
|
||||
seed = inference_args.seed
|
||||
num_videos_per_prompt = inference_args.num_videos
|
||||
seed = fastvideo_args.seed
|
||||
num_videos_per_prompt = fastvideo_args.num_videos
|
||||
|
||||
seeds = [seed + i for i in range(num_videos_per_prompt)]
|
||||
batch.seeds = seeds
|
||||
@@ -37,19 +37,19 @@ class InputValidationStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Validate and prepare inputs.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The validated batch information.
|
||||
"""
|
||||
self._generate_seeds(batch, inference_args)
|
||||
self._generate_seeds(batch, fastvideo_args)
|
||||
|
||||
# Ensure prompt is properly formatted
|
||||
if batch.prompt is None and batch.prompt_embeds is None:
|
||||
@@ -91,6 +91,6 @@ class InputValidationStage(PipelineStage):
|
||||
|
||||
# Set data type if not already set
|
||||
if batch.data_type is None:
|
||||
batch.data_type = inference_args.precision
|
||||
batch.data_type = fastvideo_args.precision
|
||||
|
||||
return batch
|
||||
|
||||
@@ -4,8 +4,9 @@ Latent preparation stage for diffusion pipelines.
|
||||
"""
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
|
||||
@@ -20,21 +21,22 @@ class LatentPreparationStage(PipelineStage):
|
||||
denoised during the diffusion process.
|
||||
"""
|
||||
|
||||
def __init__(self, scheduler) -> None:
|
||||
def __init__(self, scheduler, vae=None) -> None:
|
||||
super().__init__()
|
||||
self.scheduler = scheduler
|
||||
self.vae = vae
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Prepare initial latent variables for the diffusion process.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with prepared latent variables.
|
||||
@@ -42,7 +44,7 @@ class LatentPreparationStage(PipelineStage):
|
||||
|
||||
# Adjust video length based on VAE version if needed
|
||||
if hasattr(self, 'adjust_video_length'):
|
||||
batch = self.adjust_video_length(batch, inference_args)
|
||||
batch = self.adjust_video_length(self.vae, batch, fastvideo_args)
|
||||
# Determine batch size
|
||||
if isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
@@ -67,13 +69,16 @@ class LatentPreparationStage(PipelineStage):
|
||||
if height is None or width is None:
|
||||
raise ValueError("Height and width must be provided")
|
||||
|
||||
assert fastvideo_args.num_channels_latents is not None
|
||||
assert fastvideo_args.vae_scale_factor is not None
|
||||
|
||||
# Calculate latent shape
|
||||
shape = (
|
||||
batch_size,
|
||||
inference_args.num_channels_latents,
|
||||
fastvideo_args.num_channels_latents,
|
||||
num_frames,
|
||||
height // inference_args.vae_scale_factor,
|
||||
width // inference_args.vae_scale_factor,
|
||||
height // fastvideo_args.vae_scale_factor,
|
||||
width // fastvideo_args.vae_scale_factor,
|
||||
)
|
||||
|
||||
# Validate generator if it's a list
|
||||
@@ -101,19 +106,20 @@ class LatentPreparationStage(PipelineStage):
|
||||
|
||||
return batch
|
||||
|
||||
def adjust_video_length(self, batch: ForwardBatch,
|
||||
inference_args: InferenceArgs) -> ForwardBatch:
|
||||
def adjust_video_length(self, vae: ParallelTiledVAE, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
"""
|
||||
Adjust video length based on VAE version.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with adjusted video length.
|
||||
"""
|
||||
video_length = batch.num_frames
|
||||
temporal_scale_factor = vae.temporal_compression_ratio if vae is not None else 4
|
||||
# TODO
|
||||
batch.num_frames = (video_length - 1) // 4 + 1
|
||||
batch.num_frames = (video_length - 1) // temporal_scale_factor + 1
|
||||
return batch
|
||||
|
||||
@@ -7,8 +7,10 @@ This module contains implementations of prompt encoding stages for diffusion pip
|
||||
|
||||
from typing import TypedDict
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -59,18 +61,20 @@ class LlamaEncodingStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
|
||||
text = prompt_template_video["template"].format(batch.prompt)
|
||||
text_inputs = self.tokenizer(
|
||||
@@ -93,4 +97,34 @@ class LlamaEncodingStage(PipelineStage):
|
||||
last_hidden_state = last_hidden_state[:, crop_start:]
|
||||
batch.prompt_embeds.append(last_hidden_state)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
negative_text = prompt_template_video["template"].format(
|
||||
batch.negative_prompt)
|
||||
negative_text_inputs = self.tokenizer(
|
||||
negative_text,
|
||||
truncation=True,
|
||||
# better way to handle this?
|
||||
max_length=256,
|
||||
return_tensors="pt",
|
||||
)
|
||||
hidden_state_skip_layer = 2
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
negative_outputs = self.text_encoder(
|
||||
input_ids=negative_text_inputs["input_ids"].to(
|
||||
batch.device),
|
||||
output_hidden_states=hidden_state_skip_layer is not None,
|
||||
)
|
||||
|
||||
negative_last_hidden_state = negative_outputs.hidden_states[-(
|
||||
hidden_state_skip_layer + 1)]
|
||||
crop_start = prompt_template_video.get("crop_start", -1)
|
||||
negative_last_hidden_state = negative_last_hidden_state[:,
|
||||
crop_start:]
|
||||
assert batch.negative_prompt_embeds is not None
|
||||
batch.negative_prompt_embeds.append(negative_last_hidden_state)
|
||||
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
|
||||
@@ -0,0 +1,116 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Prompt encoding stages for diffusion pipelines.
|
||||
|
||||
This module contains implementations of prompt encoding stages for diffusion pipelines.
|
||||
"""
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class T5EncodingStage(PipelineStage):
|
||||
"""
|
||||
Stage for encoding text prompts into embeddings for diffusion models.
|
||||
|
||||
This stage handles the encoding of text prompts into the embedding space
|
||||
expected by the diffusion model.
|
||||
"""
|
||||
|
||||
def __init__(self, text_encoder, tokenizer) -> None:
|
||||
"""
|
||||
Initialize the prompt encoding stage.
|
||||
|
||||
Args:
|
||||
enable_logging: Whether to enable logging for this stage.
|
||||
is_secondary: Whether this is a secondary text encoder.
|
||||
"""
|
||||
super().__init__()
|
||||
self.text_encoder = text_encoder
|
||||
self.tokenizer = tokenizer
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Encode the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder = self.text_encoder.to(batch.device)
|
||||
|
||||
text = batch.prompt
|
||||
text_inputs = self.tokenizer(
|
||||
text,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
max_length=512,
|
||||
add_special_tokens=True,
|
||||
return_attention_mask=True,
|
||||
return_tensors="pt",
|
||||
).to(batch.device)
|
||||
text_input_ids, mask = text_inputs.input_ids, text_inputs.attention_mask
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs = self.text_encoder(
|
||||
input_ids=text_input_ids,
|
||||
attention_mask=mask,
|
||||
)
|
||||
assert torch.isnan(outputs).sum() == 0
|
||||
prompt_embeds = [u[:v] for u, v in zip(outputs, seq_lens)]
|
||||
prompt_embeds = torch.stack([
|
||||
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
|
||||
for u in prompt_embeds
|
||||
],
|
||||
dim=0)
|
||||
batch.prompt_embeds.append(prompt_embeds)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
negative_text = batch.negative_prompt
|
||||
negative_text_inputs = self.tokenizer(
|
||||
negative_text,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
max_length=512,
|
||||
add_special_tokens=True,
|
||||
return_attention_mask=True,
|
||||
return_tensors="pt",
|
||||
).to(batch.device)
|
||||
text_input_ids, mask = negative_text_inputs.input_ids, negative_text_inputs.attention_mask
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
negative_outputs = self.text_encoder(
|
||||
input_ids=text_input_ids,
|
||||
attention_mask=mask,
|
||||
)
|
||||
assert torch.isnan(negative_outputs).sum() == 0
|
||||
neg_prompt_embeds = [
|
||||
u[:v] for u, v in zip(negative_outputs, seq_lens)
|
||||
]
|
||||
neg_prompt_embeds = torch.stack([
|
||||
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
|
||||
for u in neg_prompt_embeds
|
||||
],
|
||||
dim=0)
|
||||
assert batch.negative_prompt_embeds is not None
|
||||
batch.negative_prompt_embeds.append(neg_prompt_embeds)
|
||||
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.text_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
@@ -7,7 +7,7 @@ This module contains implementations of timestep preparation stages for diffusio
|
||||
|
||||
import inspect
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
@@ -29,14 +29,14 @@ class TimestepPreparationStage(PipelineStage):
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Prepare timesteps for the diffusion process.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
inference_args: The inference arguments.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
The batch with prepared timesteps.
|
||||
|
||||
@@ -0,0 +1,79 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Wan video diffusion pipeline implementation.
|
||||
|
||||
This module contains an implementation of the Wan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (
|
||||
CLIPImageEncodingStage, ConditioningStage, DecodingStage, DenoisingStage,
|
||||
EncodingStage, InputValidationStage, LatentPreparationStage,
|
||||
T5EncodingStage, TimestepPreparationStage)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanImageToVideoPipeline(ComposedPipelineBase):
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler", \
|
||||
"image_encoder", "image_processor"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=T5EncodingStage(
|
||||
text_encoder=self.get_module("text_encoder"),
|
||||
tokenizer=self.get_module("tokenizer"),
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="image_encoding_stage",
|
||||
stage=CLIPImageEncodingStage(
|
||||
image_encoder=self.get_module("image_encoder"),
|
||||
image_processor=self.get_module("image_processor"),
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="conditioning_stage",
|
||||
stage=ConditioningStage())
|
||||
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="image_latent_preparation_stage",
|
||||
stage=EncodingStage(vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
|
||||
fastvideo_args.vae_scale_factor = vae_scale_factor
|
||||
|
||||
num_channels_latents = self.get_module("transformer").out_channels
|
||||
fastvideo_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
EntryClass = WanImageToVideoPipeline
|
||||
@@ -0,0 +1,72 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Wan video diffusion pipeline implementation.
|
||||
|
||||
This module contains an implementation of the Wan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (ConditioningStage, DecodingStage,
|
||||
DenoisingStage, InputValidationStage,
|
||||
LatentPreparationStage,
|
||||
T5EncodingStage,
|
||||
TimestepPreparationStage)
|
||||
|
||||
# TODO(will): move PRECISION_TO_TYPE to better place
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanPipeline(ComposedPipelineBase):
|
||||
|
||||
_required_config_modules = [
|
||||
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=T5EncodingStage(
|
||||
text_encoder=self.get_module("text_encoder"),
|
||||
tokenizer=self.get_module("tokenizer"),
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="conditioning_stage",
|
||||
stage=ConditioningStage())
|
||||
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
vae_scale_factor = self.get_module("vae").spatial_compression_ratio
|
||||
fastvideo_args.vae_scale_factor = vae_scale_factor
|
||||
|
||||
num_channels_latents = self.get_module("transformer").in_channels
|
||||
fastvideo_args.num_channels_latents = num_channels_latents
|
||||
|
||||
|
||||
EntryClass = WanPipeline
|
||||
@@ -18,7 +18,7 @@ def cuda_platform_plugin() -> Optional[str]:
|
||||
|
||||
try:
|
||||
from fastvideo.v1.utils import import_pynvml
|
||||
pynvml = import_pynvml()
|
||||
pynvml = import_pynvml() # type: ignore[no-untyped-call]
|
||||
pynvml.nvmlInit()
|
||||
try:
|
||||
# NOTE: Edge case: vllm cpu build on a GPU machine.
|
||||
|
||||
@@ -26,7 +26,7 @@ logger = init_logger(__name__)
|
||||
_P = ParamSpec("_P")
|
||||
_R = TypeVar("_R")
|
||||
|
||||
pynvml = import_pynvml()
|
||||
pynvml = import_pynvml() # type: ignore[no-untyped-call]
|
||||
|
||||
# pytorch 2.5 uses cudnn sdpa by default, which will cause crash on some models
|
||||
# see https://github.com/huggingface/diffusers/issues/9704 for details
|
||||
@@ -110,14 +110,13 @@ class CudaPlatformBase(Platform):
|
||||
return float(torch.cuda.max_memory_allocated(device))
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend, head_size, dtype,
|
||||
distributed) -> str:
|
||||
def get_attn_backend_cls(cls, selected_backend: Optional[_Backend],
|
||||
head_size: int, dtype: torch.dtype) -> str:
|
||||
# TODO(will): maybe come up with a more general interface for local attention
|
||||
# if distributed is False, we always try to use Flash attn
|
||||
|
||||
logger.info(
|
||||
"Distributed attention=%s, trying FASTVIDEO_ATTENTION_BACKEND=%s",
|
||||
distributed, envs.FASTVIDEO_ATTENTION_BACKEND)
|
||||
logger.info("Trying FASTVIDEO_ATTENTION_BACKEND=%s",
|
||||
envs.FASTVIDEO_ATTENTION_BACKEND)
|
||||
if selected_backend == _Backend.SLIDING_TILE_ATTN:
|
||||
try:
|
||||
from st_attn import sliding_tile_attention # noqa: F401
|
||||
@@ -132,10 +131,10 @@ class CudaPlatformBase(Platform):
|
||||
logger.info(
|
||||
"Sliding Tile Attention backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == _Backend.FLASH_ATTN:
|
||||
pass
|
||||
elif selected_backend == _Backend.TORCH_SDPA:
|
||||
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
|
||||
elif selected_backend == _Backend.FLASH_ATTN or selected_backend is None:
|
||||
pass
|
||||
elif selected_backend:
|
||||
raise ValueError(f"Invalid attention backend for {cls.device_name}")
|
||||
|
||||
|
||||
@@ -87,8 +87,8 @@ class Platform:
|
||||
return self._enum == PlatformEnum.CUDA
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend: _Backend, head_size: int,
|
||||
dtype: torch.dtype, distributed: bool) -> str:
|
||||
def get_attn_backend_cls(cls, selected_backend: Optional[_Backend],
|
||||
head_size: int, dtype: torch.dtype) -> str:
|
||||
"""Get the attention backend class of a device."""
|
||||
return ""
|
||||
|
||||
|
||||
@@ -11,12 +11,12 @@ from einops import rearrange
|
||||
|
||||
from fastvideo.v1.distributed import (init_distributed_environment,
|
||||
initialize_model_parallel)
|
||||
from fastvideo.v1.inference_args import InferenceArgs, prepare_inference_args
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, prepare_fastvideo_args
|
||||
# Fix the import path
|
||||
from fastvideo.v1.inference_engine import InferenceEngine
|
||||
|
||||
|
||||
def initialize_distributed_and_parallelism(inference_args: InferenceArgs):
|
||||
def initialize_distributed_and_parallelism(fastvideo_args: FastVideoArgs):
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
@@ -25,29 +25,33 @@ def initialize_distributed_and_parallelism(inference_args: InferenceArgs):
|
||||
rank=rank,
|
||||
local_rank=local_rank)
|
||||
device_str = f"cuda:{local_rank}"
|
||||
inference_args.device_str = device_str
|
||||
inference_args.device = torch.device(device_str)
|
||||
fastvideo_args.device_str = device_str
|
||||
fastvideo_args.device = torch.device(device_str)
|
||||
assert fastvideo_args.sp_size is not None
|
||||
assert fastvideo_args.tp_size is not None
|
||||
initialize_model_parallel(
|
||||
sequence_model_parallel_size=inference_args.sp_size,
|
||||
tensor_model_parallel_size=inference_args.tp_size,
|
||||
sequence_model_parallel_size=fastvideo_args.sp_size,
|
||||
tensor_model_parallel_size=fastvideo_args.tp_size,
|
||||
)
|
||||
|
||||
|
||||
def main(inference_args: InferenceArgs):
|
||||
initialize_distributed_and_parallelism(inference_args)
|
||||
engine = InferenceEngine.create_engine(inference_args, )
|
||||
def main(fastvideo_args: FastVideoArgs):
|
||||
initialize_distributed_and_parallelism(fastvideo_args)
|
||||
engine = InferenceEngine.create_engine(fastvideo_args, )
|
||||
|
||||
if inference_args.prompt_path is not None:
|
||||
with open(inference_args.prompt_path) as f:
|
||||
if fastvideo_args.prompt_path is not None:
|
||||
with open(fastvideo_args.prompt_path) as f:
|
||||
prompts = [line.strip() for line in f.readlines()]
|
||||
else:
|
||||
prompts = [inference_args.prompt]
|
||||
if fastvideo_args.prompt is None:
|
||||
raise ValueError("prompt or prompt_path is required")
|
||||
prompts = [fastvideo_args.prompt]
|
||||
|
||||
# Process each prompt
|
||||
for prompt in prompts:
|
||||
outputs = engine.run(
|
||||
prompt=prompt,
|
||||
inference_args=inference_args,
|
||||
fastvideo_args=fastvideo_args,
|
||||
)
|
||||
|
||||
# Process outputs
|
||||
@@ -59,13 +63,13 @@ def main(inference_args: InferenceArgs):
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
|
||||
# Save video
|
||||
os.makedirs(os.path.dirname(inference_args.output_path), exist_ok=True)
|
||||
imageio.mimsave(os.path.join(inference_args.output_path,
|
||||
os.makedirs(os.path.dirname(fastvideo_args.output_path), exist_ok=True)
|
||||
imageio.mimsave(os.path.join(fastvideo_args.output_path,
|
||||
f"{prompt[:100]}.mp4"),
|
||||
frames,
|
||||
fps=inference_args.fps)
|
||||
fps=fastvideo_args.fps)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
inference_args = prepare_inference_args(sys.argv[1:])
|
||||
main(inference_args)
|
||||
fastvideo_args = prepare_fastvideo_args(sys.argv[1:])
|
||||
main(fastvideo_args)
|
||||
|
||||
@@ -1,9 +1,14 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import pytest
|
||||
import torch.distributed as dist
|
||||
|
||||
from fastvideo.v1.distributed import (destroy_model_parallel,
|
||||
init_distributed_environment,
|
||||
initialize_model_parallel)
|
||||
import pytest
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.v1.distributed import (init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
cleanup_dist_env_and_memory)
|
||||
|
||||
|
||||
@pytest.fixture(scope="function")
|
||||
@@ -13,6 +18,9 @@ def distributed_setup():
|
||||
|
||||
This ensures proper cleanup even if tests fail.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
np.random.seed(42)
|
||||
|
||||
init_distributed_environment(world_size=1,
|
||||
rank=0,
|
||||
distributed_init_method="env://",
|
||||
@@ -23,6 +31,4 @@ def distributed_setup():
|
||||
backend="nccl")
|
||||
yield
|
||||
|
||||
if dist.is_initialized():
|
||||
destroy_model_parallel()
|
||||
dist.destroy_process_group()
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# TODO: check if correct
|
||||
import os
|
||||
|
||||
@@ -10,7 +11,7 @@ from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
|
||||
load_tokenizer)
|
||||
# from fastvideo.v1.models.hunyuan.text_encoder import load_text_encoder, load_tokenizer
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
|
||||
@@ -19,10 +20,6 @@ logger = init_logger(__name__)
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
# Set fixed random seed for reproducibility
|
||||
torch.manual_seed(42)
|
||||
np.random.seed(42)
|
||||
|
||||
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
@@ -41,7 +38,7 @@ def test_clip_encoder():
|
||||
- Load models with the same weights and parameters
|
||||
- Produce nearly identical outputs for the same input prompts
|
||||
"""
|
||||
args = InferenceArgs(model_path="openai/clip-vit-large-patch14",
|
||||
args = FastVideoArgs(model_path="openai/clip-vit-large-patch14",
|
||||
precision="float16")
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
@@ -8,7 +9,7 @@ from transformers import AutoConfig
|
||||
from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
|
||||
load_tokenizer)
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
@@ -18,10 +19,6 @@ logger = init_logger(__name__)
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
# Set fixed random seed for reproducibility
|
||||
torch.manual_seed(42)
|
||||
np.random.seed(42)
|
||||
|
||||
BASE_MODEL_PATH = "hunyuanvideo-community/HunyuanVideo"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
@@ -41,7 +38,7 @@ def test_llama_encoder():
|
||||
- Load models with the same weights and parameters
|
||||
- Produce nearly identical outputs for the same input prompts
|
||||
"""
|
||||
args = InferenceArgs(model_path="meta-llama/Llama-2-7b-hf",
|
||||
args = FastVideoArgs(model_path="meta-llama/Llama-2-7b-hf",
|
||||
precision="float16")
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
@@ -0,0 +1,136 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from transformers import AutoConfig, AutoTokenizer, UMT5EncoderModel
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
os.environ["MASTER_PORT"] = "29503"
|
||||
|
||||
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
|
||||
local_dir=os.path.join(
|
||||
'data', BASE_MODEL_PATH))
|
||||
TEXT_ENCODER_PATH = os.path.join(MODEL_PATH, "text_encoder")
|
||||
TOKENIZER_PATH = os.path.join(MODEL_PATH, "tokenizer")
|
||||
|
||||
|
||||
@pytest.mark.usefixtures("distributed_setup")
|
||||
def test_t5_encoder():
|
||||
# Initialize the two model implementations
|
||||
hf_config = AutoConfig.from_pretrained(TEXT_ENCODER_PATH)
|
||||
print(hf_config)
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
precision = torch.float16
|
||||
model1 = UMT5EncoderModel.from_pretrained(TEXT_ENCODER_PATH).to(
|
||||
precision).to(device).eval()
|
||||
tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
|
||||
|
||||
loader = TextEncoderLoader()
|
||||
model2 = loader.load_model(TEXT_ENCODER_PATH, hf_config, device)
|
||||
|
||||
# Convert to float16 and move to device
|
||||
model2 = model2.to(precision)
|
||||
model2 = model2.to(device)
|
||||
model2.eval()
|
||||
|
||||
# Sanity check weights between the two models
|
||||
logger.info("Comparing model weights for sanity check...")
|
||||
params1 = dict(model1.named_parameters())
|
||||
params2 = dict(model2.named_parameters())
|
||||
|
||||
# Check number of parameters
|
||||
logger.info("Model1 has %s parameters", len(params1))
|
||||
logger.info("Model2 has %s parameters", len(params2))
|
||||
|
||||
weight_diffs = []
|
||||
# check if embed_tokens are the same
|
||||
weights = ["encoder.block.{}.layer.0.layer_norm.weight", "encoder.block.{}.layer.0.SelfAttention.relative_attention_bias.weight", \
|
||||
"encoder.block.{}.layer.0.SelfAttention.o.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_0.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_1.weight",\
|
||||
"encoder.block.{}.layer.1.DenseReluDense.wo.weight", \
|
||||
"encoder.block.{}.layer.1.layer_norm.weight", "encoder.final_layer_norm.weight", "shared.weight"]
|
||||
for idx in range(hf_config.num_hidden_layers):
|
||||
for w in weights:
|
||||
name1 = w.format(idx)
|
||||
name2 = w.format(idx)
|
||||
p1 = params1[name1]
|
||||
p2 = params2[name2]
|
||||
assert p1.dtype == p2.dtype
|
||||
try:
|
||||
logger.info("Parameter: %s vs %s", name1, name2)
|
||||
max_diff = torch.max(torch.abs(p1 - p2)).item()
|
||||
mean_diff = torch.mean(torch.abs(p1 - p2)).item()
|
||||
weight_diffs.append((name1, name2, max_diff, mean_diff))
|
||||
logger.info(" Max diff: %s, Mean diff: %s", max_diff,
|
||||
mean_diff)
|
||||
except Exception as e:
|
||||
logger.info("Error comparing %s and %s: %s", name1, name2, e)
|
||||
|
||||
# Test with some sample prompts
|
||||
prompts = [
|
||||
"Once upon a time", "The quick brown fox jumps over",
|
||||
"In a galaxy far, far away"
|
||||
]
|
||||
|
||||
logger.info("Testing T5 encoder with sample prompts")
|
||||
|
||||
with torch.no_grad():
|
||||
for prompt in prompts:
|
||||
logger.info("Testing prompt: %s", prompt)
|
||||
|
||||
# Tokenize the prompt
|
||||
tokens = tokenizer(prompt,
|
||||
padding="max_length",
|
||||
max_length=512,
|
||||
truncation=True,
|
||||
return_tensors="pt").to(device)
|
||||
|
||||
# Get outputs from HuggingFace implementation
|
||||
# filter out padding input_ids
|
||||
# tokens.input_ids = tokens.input_ids[tokens.attention_mask==1]
|
||||
# tokens.attention_mask = tokens.attention_mask[tokens.attention_mask==1]
|
||||
outputs1 = model1(input_ids=tokens.input_ids,
|
||||
attention_mask=tokens.attention_mask,
|
||||
output_hidden_states=True).last_hidden_state
|
||||
print("--------------------------------")
|
||||
logger.info("Testing model2")
|
||||
|
||||
# Get outputs from our implementation
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs2 = model2(
|
||||
input_ids=tokens.input_ids,
|
||||
attention_mask=tokens.attention_mask,
|
||||
)
|
||||
|
||||
# Compare last hidden states
|
||||
last_hidden_state1 = outputs1[tokens.attention_mask == 1]
|
||||
last_hidden_state2 = outputs2[tokens.attention_mask == 1]
|
||||
|
||||
assert last_hidden_state1.shape == last_hidden_state2.shape, \
|
||||
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
|
||||
|
||||
max_diff_hidden = torch.max(
|
||||
torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
mean_diff_hidden = torch.mean(
|
||||
torch.abs(last_hidden_state1 - last_hidden_state2))
|
||||
|
||||
logger.info("Maximum difference in last hidden states: %s",
|
||||
max_diff_hidden.item())
|
||||
logger.info("Mean difference in last hidden states: %s",
|
||||
mean_diff_hidden.item())
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
assert mean_diff_hidden < 1e-4, \
|
||||
f"Hidden states differ significantly: mean diff = {mean_diff_hidden.item()}"
|
||||
assert max_diff_hidden < 1e-4, \
|
||||
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
|
||||
@@ -1,179 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from diffusers import AutoencoderKLHunyuanVideo as DiffusersHunyuanVAE
|
||||
from safetensors.torch import load_file
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.hunyuanvae import (
|
||||
AutoencoderKLHunyuanVideo as MyHunyuanVAE)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def initialize_identical_weights(model1, model2, seed=42):
|
||||
"""Initialize both models with identical weights using a fixed seed for reproducibility."""
|
||||
# Get all parameters from both models
|
||||
params1 = dict(model1.named_parameters())
|
||||
params2 = dict(model2.named_parameters())
|
||||
|
||||
# Initialize each layer with identical values
|
||||
with torch.no_grad():
|
||||
# Initialize weights
|
||||
for name1, param1 in params1.items():
|
||||
if 'weight' in name1:
|
||||
# Set seed before each weight initialization
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param1, mean=0.0, std=0.05)
|
||||
|
||||
for name2, param2 in params2.items():
|
||||
if 'weight' in name2:
|
||||
# Reset seed to get same initialization
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param2, mean=0.0, std=0.05)
|
||||
|
||||
# Initialize biases
|
||||
for name1, param1 in params1.items():
|
||||
if 'bias' in name1:
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param1, mean=0.0, std=0.05)
|
||||
param1.data = param1.data.to(torch.bfloat16)
|
||||
|
||||
for name2, param2 in params2.items():
|
||||
if 'bias' in name2:
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param2, mean=0.0, std=0.05)
|
||||
param2.data = param2.data.to(torch.bfloat16)
|
||||
|
||||
logger.info("Both models initialized with identical weights in bfloat16")
|
||||
return model1, model2
|
||||
|
||||
|
||||
def setup_args():
|
||||
parser = argparse.ArgumentParser(description='HunyuanVAE Test')
|
||||
parser.add_argument('--in-channels',
|
||||
type=int,
|
||||
default=4,
|
||||
help='Number of input channels')
|
||||
parser.add_argument('--out-channels',
|
||||
type=int,
|
||||
default=4,
|
||||
help='Number of output channels')
|
||||
parser.add_argument('--latent-channels',
|
||||
type=int,
|
||||
default=4,
|
||||
help='Number of latent channels')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def test_hunyuan_vae():
|
||||
args = setup_args()
|
||||
|
||||
# Set fixed random seed for reproducibility
|
||||
torch.manual_seed(42)
|
||||
np.random.seed(42)
|
||||
|
||||
# Model parameters
|
||||
in_channels = args.in_channels
|
||||
out_channels = args.out_channels
|
||||
latent_channels = args.latent_channels
|
||||
|
||||
device = torch.device("cuda:0")
|
||||
print
|
||||
# Initialize the two model implementations
|
||||
path = "data/hunyuanvideo-community/HunyuanVideo/vae"
|
||||
config_path = os.path.join(path, "config.json")
|
||||
config = json.load(open(config_path))
|
||||
config.pop("_class_name")
|
||||
config.pop("_diffusers_version")
|
||||
model1 = MyHunyuanVAE(**config).to(torch.bfloat16)
|
||||
|
||||
model2 = DiffusersHunyuanVAE(**config).to(torch.bfloat16)
|
||||
|
||||
loaded = load_file(os.path.join(path,
|
||||
"diffusion_pytorch_model.safetensors"))
|
||||
model1.load_state_dict(loaded)
|
||||
model2.load_state_dict(loaded)
|
||||
|
||||
# Set both models to eval mode
|
||||
model1.eval()
|
||||
model2.eval()
|
||||
|
||||
# Move to GPU
|
||||
model1 = model1.to(device)
|
||||
model2 = model2.to(device)
|
||||
|
||||
model1.enable_tiling(tile_sample_min_height=32,
|
||||
tile_sample_min_width=32,
|
||||
tile_sample_min_num_frames=8,
|
||||
tile_sample_stride_height=16,
|
||||
tile_sample_stride_width=16,
|
||||
tile_sample_stride_num_frames=4)
|
||||
model2.enable_tiling(tile_sample_min_height=32,
|
||||
tile_sample_min_width=32,
|
||||
tile_sample_min_num_frames=8,
|
||||
tile_sample_stride_height=16,
|
||||
tile_sample_stride_width=16,
|
||||
tile_sample_stride_num_frames=4)
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
|
||||
# Video input [B, C, T, H, W]
|
||||
input_tensor = torch.randn(batch_size,
|
||||
3,
|
||||
21,
|
||||
64,
|
||||
64,
|
||||
device=device,
|
||||
dtype=torch.bfloat16)
|
||||
|
||||
# Disable gradients for inference
|
||||
with torch.no_grad():
|
||||
# Test encoding
|
||||
logger.info("Testing encoding...")
|
||||
latent1 = model1.encode(input_tensor).mean
|
||||
print("--------------------------------")
|
||||
latent2 = model2.encode(input_tensor).latent_dist.mean
|
||||
# Check if latents have the same shape
|
||||
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
|
||||
assert latent1.shape == latent2.shape, f"Latent shapes don't match: {latent1.shape} vs {latent2.shape}"
|
||||
# Check if latents are similar
|
||||
max_diff_encode = torch.max(torch.abs(latent1 - latent2))
|
||||
mean_diff_encode = torch.mean(torch.abs(latent1 - latent2))
|
||||
logger.info(
|
||||
f"Maximum difference between encoded latents: {max_diff_encode.item()}"
|
||||
)
|
||||
logger.info(
|
||||
f"Mean difference between encoded latents: {mean_diff_encode.item()}"
|
||||
)
|
||||
assert max_diff_encode < 1e-4, f"Encoded latents differ significantly: max diff = {max_diff_encode.item()}"
|
||||
# Test decoding
|
||||
logger.info("Testing decoding...")
|
||||
output1 = model1.decode(latent1)
|
||||
output2 = model2.decode(latent2).sample
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
|
||||
# Check if outputs are similar
|
||||
max_diff_decode = torch.max(torch.abs(output1 - output2))
|
||||
mean_diff_decode = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info(
|
||||
f"Maximum difference between decoded outputs: {max_diff_decode.item()}"
|
||||
)
|
||||
logger.info(
|
||||
f"Mean difference between decoded outputs: {mean_diff_decode.item()}"
|
||||
)
|
||||
assert max_diff_decode < 1e-4, f"Decoded outputs differ significantly: max diff = {max_diff_decode.item()}"
|
||||
|
||||
logger.info(
|
||||
"Test passed! Both VAE implementations produce similar outputs.")
|
||||
logger.info("Test completed successfully")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_hunyuan_vae()
|
||||
@@ -1,260 +0,0 @@
|
||||
import argparse
|
||||
import os
|
||||
from itertools import chain
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from torch.distributed.device_mesh import init_device_mesh
|
||||
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, destroy_distributed_environment,
|
||||
destroy_model_parallel, get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size, init_distributed_environment,
|
||||
initialize_model_parallel)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.dits.hunyuanvideo import (
|
||||
HunyuanVideoTransformer3DModel as HunyuanVideoDit)
|
||||
from fastvideo.v1.models.hunyuan.modules.models import (
|
||||
HYVideoDiffusionTransformer)
|
||||
from fastvideo.v1.models.loader.fsdp_load import shard_model
|
||||
from fastvideo.v1.utils.parallel_states import (
|
||||
initialize_sequence_parallel_state)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def initialize_identical_weights(model1, model2, seed=42):
|
||||
"""Initialize both models with identical weights using a fixed seed for reproducibility."""
|
||||
# Get all parameters from both models
|
||||
params1 = dict(model1.named_parameters())
|
||||
params2 = dict(model2.named_parameters())
|
||||
|
||||
# Initialize each layer with identical values
|
||||
with torch.no_grad():
|
||||
# Initialize weights
|
||||
for name1, param1 in params1.items():
|
||||
if 'weight' in name1:
|
||||
# Set seed before each weight initialization
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param1, mean=0.0, std=0.05)
|
||||
|
||||
for name2, param2 in params2.items():
|
||||
if 'weight' in name2:
|
||||
# Reset seed to get same initialization
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param2, mean=0.0, std=0.05)
|
||||
|
||||
# Initialize biases
|
||||
for name1, param1 in params1.items():
|
||||
if 'bias' in name1:
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param1, mean=0.0, std=0.05)
|
||||
param1.data = param1.data.to(torch.bfloat16)
|
||||
|
||||
for name2, param2 in params2.items():
|
||||
if 'bias' in name2:
|
||||
torch.manual_seed(seed)
|
||||
nn.init.normal_(param2, mean=0.0, std=0.05)
|
||||
param2.data = param2.data.to(torch.bfloat16)
|
||||
|
||||
logger.info("Both models initialized with identical weights in bfloat16")
|
||||
return model1, model2
|
||||
|
||||
|
||||
def setup_args():
|
||||
parser = argparse.ArgumentParser(
|
||||
description='Distributed HunyuanVideo Test')
|
||||
parser.add_argument('--sequence_model_parallel_size',
|
||||
type=int,
|
||||
default=1,
|
||||
help='Degree of sequence model parallelism')
|
||||
parser.add_argument('--hidden-size',
|
||||
type=int,
|
||||
default=128,
|
||||
help='Hidden size for the model')
|
||||
parser.add_argument('--heads-num',
|
||||
type=int,
|
||||
default=4,
|
||||
help='Number of attention heads')
|
||||
parser.add_argument('--double-blocks-depth',
|
||||
type=int,
|
||||
default=2,
|
||||
help='Number of double stream blocks')
|
||||
parser.add_argument('--single-blocks-depth',
|
||||
type=int,
|
||||
default=2,
|
||||
help='Number of single stream blocks')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def test_hunyuanvideo_distributed():
|
||||
args = setup_args()
|
||||
|
||||
# Initialize distributed environment
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
|
||||
logger.info(
|
||||
f"Initializing process: rank={rank}, local_rank={local_rank}, world_size={world_size}"
|
||||
)
|
||||
|
||||
# Initialize distributed environment
|
||||
init_distributed_environment(world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=local_rank)
|
||||
|
||||
# Initialize tensor model parallel groups
|
||||
initialize_model_parallel(
|
||||
sequence_model_parallel_size=args.sequence_model_parallel_size)
|
||||
initialize_sequence_parallel_state(world_size)
|
||||
# Get tensor parallel info
|
||||
sp_rank = get_sequence_model_parallel_rank()
|
||||
sp_world_size = get_sequence_model_parallel_world_size()
|
||||
|
||||
logger.info(
|
||||
f"Process rank {rank} initialized with SP rank {sp_rank} in SP world size {sp_world_size}"
|
||||
)
|
||||
|
||||
# Set fixed random seed for reproducibility
|
||||
torch.manual_seed(42)
|
||||
np.random.seed(42)
|
||||
|
||||
# Small model parameters for testing
|
||||
hidden_size = args.hidden_size
|
||||
heads_num = args.heads_num
|
||||
mm_double_blocks_depth = args.double_blocks_depth
|
||||
mm_single_blocks_depth = args.single_blocks_depth
|
||||
patch_size = [1, 2, 2]
|
||||
torch.cuda.set_device(f"cuda:{local_rank}")
|
||||
# Initialize the two model implementations
|
||||
model1 = HunyuanVideoDit(
|
||||
patch_size=2,
|
||||
patch_size_t=1,
|
||||
in_channels=4,
|
||||
out_channels=4,
|
||||
attention_head_dim=hidden_size // heads_num,
|
||||
num_attention_heads=heads_num,
|
||||
num_layers=mm_double_blocks_depth,
|
||||
num_single_layers=mm_single_blocks_depth,
|
||||
rope_axes_dim=[8, 16, 8], # sum = hidden_size // heads_num = 32
|
||||
dtype=torch.bfloat16).to(torch.bfloat16)
|
||||
model2 = HYVideoDiffusionTransformer(
|
||||
patch_size=patch_size,
|
||||
in_channels=4,
|
||||
hidden_size=hidden_size,
|
||||
heads_num=heads_num,
|
||||
mm_double_blocks_depth=mm_double_blocks_depth,
|
||||
mm_single_blocks_depth=mm_single_blocks_depth,
|
||||
rope_dim_list=[8, 16, 8], # sum = hidden_size // heads_num = 32
|
||||
dtype=torch.bfloat16).to(torch.bfloat16)
|
||||
|
||||
# print("--------------------------------")
|
||||
# for name, param in model3.named_parameters():
|
||||
# print(name)
|
||||
# import pdb; pdb.set_trace()
|
||||
# # Initialize with identical weights
|
||||
model1, model2 = initialize_identical_weights(model1, model2, seed=42)
|
||||
device_mesh = init_device_mesh(
|
||||
"cuda",
|
||||
mesh_shape=(sp_world_size, ),
|
||||
mesh_dim_names=("dp", ),
|
||||
)
|
||||
shard_model(model1, cpu_offload=False, reshard_after_forward=True)
|
||||
for n, p in chain(model1.named_parameters(), model1.named_buffers()):
|
||||
if p.is_meta:
|
||||
raise RuntimeError(
|
||||
f"Unexpected param or buffer {n} on meta device.")
|
||||
for p in model1.parameters():
|
||||
p.requires_grad = False
|
||||
# Set both models to eval mode
|
||||
model1.eval()
|
||||
model2.eval()
|
||||
|
||||
# Move to GPU based on local rank (0 or 1 for 2 GPUs)
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
model1 = model1.to(device)
|
||||
model2 = model2.to(device)
|
||||
|
||||
# Create identical inputs for both models
|
||||
batch_size = 1
|
||||
seq_len = 3
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size,
|
||||
4,
|
||||
8,
|
||||
16,
|
||||
16,
|
||||
device=device,
|
||||
dtype=torch.bfloat16)
|
||||
chunk_per_rank = hidden_states.shape[2] // sp_world_size
|
||||
hidden_states = hidden_states[:, :, sp_rank * chunk_per_rank:(sp_rank + 1) *
|
||||
chunk_per_rank]
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size,
|
||||
seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=torch.bfloat16)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.tensor([500], device=device, dtype=torch.bfloat16)
|
||||
|
||||
# Attention mask for text
|
||||
encoder_attention_mask = torch.ones(batch_size,
|
||||
seq_len,
|
||||
device=device,
|
||||
dtype=torch.bfloat16)
|
||||
guidance = torch.tensor([1.0], device=device, dtype=torch.bfloat16)
|
||||
# Disable gradients for inference
|
||||
with torch.no_grad():
|
||||
output1 = model1(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
)
|
||||
print("--------------------------------")
|
||||
output2, _ = model2(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
)
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
|
||||
# Check if outputs are similar (allowing for small numerical differences)
|
||||
max_diff = torch.max(torch.abs(output1 - output2))
|
||||
logger.info(f"Maximum difference between outputs: {max_diff.item()}")
|
||||
# mean diff
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
logger.info(f"Mean difference between outputs: {mean_diff.item()}")
|
||||
# diff sum
|
||||
diff_sum = torch.sum(torch.abs(output1 - output2))
|
||||
logger.info(f"Diff sum between outputs: {diff_sum.item()}")
|
||||
# sum
|
||||
sum_output1 = torch.sum(output1.float())
|
||||
sum_output2 = torch.sum(output2.float())
|
||||
logger.info(f"Rank {sp_rank} Sum of output1: {sum_output1.item()}")
|
||||
logger.info(f"Rank {sp_rank} Sum of output2: {sum_output2.item()}")
|
||||
# The outputs should be very close if not identical
|
||||
assert max_diff < 1e-3, f"Outputs differ significantly: max diff = {max_diff.item()}" # Increased tolerance for bf16
|
||||
|
||||
logger.info(
|
||||
"Test passed! Both model implementations produce the same outputs.")
|
||||
|
||||
# Clean up
|
||||
logger.info("Cleaning up distributed environment")
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
logger.info("Test completed successfully")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_hunyuanvideo_distributed()
|
||||
@@ -1,212 +0,0 @@
|
||||
import argparse
|
||||
import glob
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.models.hunyuan.modules.models import (
|
||||
HUNYUAN_VIDEO_CONFIG, HYVideoDiffusionTransformer)
|
||||
from fastvideo.utils.parallel_states import initialize_sequence_parallel_state
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, destroy_distributed_environment,
|
||||
destroy_model_parallel, get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size, init_distributed_environment,
|
||||
initialize_model_parallel)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.dits.hunyuanvideo import (
|
||||
HunyuanVideoTransformer3DModel as HunyuanVideoDit)
|
||||
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def setup_args():
|
||||
parser = argparse.ArgumentParser(
|
||||
description='Distributed HunyuanVideo Test')
|
||||
parser.add_argument('--sequence_model_parallel_size',
|
||||
type=int,
|
||||
default=1,
|
||||
help='Degree of sequence model parallelism')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def test_hunyuanvideo_distributed():
|
||||
args = setup_args()
|
||||
|
||||
# Initialize distributed environment
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
|
||||
logger.info(
|
||||
f"Initializing process: rank={rank}, local_rank={local_rank}, world_size={world_size}"
|
||||
)
|
||||
|
||||
# Initialize distributed environment
|
||||
init_distributed_environment(world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=local_rank)
|
||||
torch.cuda.set_device(f"cuda:{local_rank}")
|
||||
# Initialize tensor model parallel groups
|
||||
initialize_model_parallel(
|
||||
sequence_model_parallel_size=args.sequence_model_parallel_size)
|
||||
initialize_sequence_parallel_state(args.sequence_model_parallel_size)
|
||||
# Get tensor parallel info
|
||||
sp_rank = get_sequence_model_parallel_rank()
|
||||
sp_world_size = get_sequence_model_parallel_world_size()
|
||||
|
||||
logger.info(
|
||||
f"Process rank {rank} initialized with SP rank {sp_rank} in SP world size {sp_world_size}"
|
||||
)
|
||||
|
||||
# load data/hunyuanvideo_community/transformer/config.json
|
||||
with open(
|
||||
"data/hunyuanvideo-community/HunyuanVideo/transformer/config.json") as f:
|
||||
config = json.load(f)
|
||||
# remove "_class_name": "HunyuanVideoTransformer3DModel", "_diffusers_version": "0.32.0.dev0",
|
||||
# TODO: write normalize config function
|
||||
config.pop("_class_name")
|
||||
config.pop("_diffusers_version")
|
||||
# load data/hunyuanvideo_community/transformer/*.safetensors
|
||||
weight_dir_list = glob.glob(
|
||||
"data/hunyuanvideo-community/HunyuanVideo/transformer/*.safetensors")
|
||||
# to str
|
||||
weight_dir_list = [str(path) for path in weight_dir_list]
|
||||
model1 = load_fsdp_model(HunyuanVideoDit,
|
||||
init_params=config,
|
||||
weight_dir_list=weight_dir_list,
|
||||
device=torch.device(f"cuda:{local_rank}"),
|
||||
cpu_offload=False)
|
||||
|
||||
# successfully sharded the model (hunyuanvideo bf16 should take around 26GB in total)
|
||||
total_params = sum(p.numel() for p in model1.parameters())
|
||||
logger.info(f"Total parameters: {total_params / 1e9}B")
|
||||
|
||||
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
|
||||
weight_sum_model1 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model1.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model1 = weight_sum_model1 / total_params
|
||||
|
||||
model2 = HYVideoDiffusionTransformer(
|
||||
in_channels=16,
|
||||
out_channels=16,
|
||||
**HUNYUAN_VIDEO_CONFIG["HYVideo-T/2-cfgdistill"],
|
||||
device=torch.device(f"cuda:{local_rank}"),
|
||||
dtype=torch.bfloat16).bfloat16()
|
||||
# data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt
|
||||
state_dict = torch.load(
|
||||
"/mbz/users/hao.zhang/peiyuan/FastVideo/data/hunyuan/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
|
||||
map_location=lambda storage, loc: storage)["module"]
|
||||
model2.load_state_dict(state_dict, strict=True)
|
||||
model2.to(torch.device(f"cuda:{local_rank}")).bfloat16()
|
||||
print("load state dict done")
|
||||
|
||||
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
|
||||
total_params_model2 = sum(p.numel() for p in model2.parameters())
|
||||
weight_sum_model2 = sum(
|
||||
p.to(torch.float64).sum().item() for p in model2.parameters())
|
||||
# Also calculate mean for more stable comparison
|
||||
weight_mean_model2 = weight_sum_model2 / total_params_model2
|
||||
logger.info(f"Model 2 weight sum: {weight_sum_model2}")
|
||||
logger.info(f"Model 2 weight mean: {weight_mean_model2}")
|
||||
|
||||
# Set both models to eval mode
|
||||
model1.eval()
|
||||
model2.eval()
|
||||
|
||||
# Create random inputs for testing
|
||||
batch_size = 1
|
||||
seq_len = 3
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
|
||||
# Video latents [B, C, T, H, W]
|
||||
hidden_states = torch.randn(batch_size,
|
||||
16,
|
||||
8,
|
||||
16,
|
||||
16,
|
||||
device=device,
|
||||
dtype=torch.bfloat16)
|
||||
chunk_per_rank = hidden_states.shape[2] // sp_world_size
|
||||
hidden_states = hidden_states[:, :, sp_rank * chunk_per_rank:(sp_rank + 1) *
|
||||
chunk_per_rank]
|
||||
|
||||
# Text embeddings [B, L, D] (including global token)
|
||||
encoder_hidden_states = torch.randn(batch_size,
|
||||
seq_len + 1,
|
||||
4096,
|
||||
device=device,
|
||||
dtype=torch.bfloat16)
|
||||
|
||||
# Timestep
|
||||
timestep = torch.tensor([500], device=device, dtype=torch.bfloat16)
|
||||
|
||||
# Attention mask for text
|
||||
encoder_attention_mask = torch.ones(batch_size,
|
||||
seq_len,
|
||||
device=device,
|
||||
dtype=torch.bfloat16)
|
||||
guidance = torch.tensor([1.0], device=device, dtype=torch.bfloat16)
|
||||
|
||||
# Disable gradients for inference
|
||||
with torch.no_grad():
|
||||
# Run inference on model1
|
||||
with torch.amp.autocast(device_type="cuda", dtype=torch.bfloat16):
|
||||
logger.info("Running inference on model1")
|
||||
output1 = model1(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
)
|
||||
logger.info("Model 1 inference completed")
|
||||
|
||||
# Run inference on model2
|
||||
output2, _ = model2(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
)
|
||||
logger.info("Model 2 inference completed")
|
||||
|
||||
# Check if outputs have the same shape
|
||||
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
|
||||
|
||||
# Compare weight sums and means
|
||||
logger.info(f"Model 1 weight sum: {weight_sum_model1}")
|
||||
logger.info(f"Model 2 weight sum: {weight_sum_model2}")
|
||||
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
|
||||
logger.info(f"Weight sum difference: {weight_sum_diff}")
|
||||
|
||||
logger.info(f"Model 1 weight mean: {weight_mean_model1}")
|
||||
logger.info(f"Model 2 weight mean: {weight_mean_model2}")
|
||||
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
|
||||
logger.info(f"Weight mean difference: {weight_mean_diff}")
|
||||
|
||||
# mean diff
|
||||
mean_diff = torch.mean(torch.abs(output1 - output2))
|
||||
assert mean_diff < 1e-2, f"Mean difference between outputs: {mean_diff.item()}"
|
||||
|
||||
# diff sum
|
||||
diff_sum = torch.sum(torch.abs(output1 - output2))
|
||||
logger.info(f"Diff sum between outputs: {diff_sum.item()}")
|
||||
|
||||
# sum
|
||||
sum_output1 = torch.sum(output1.float())
|
||||
sum_output2 = torch.sum(output2.float())
|
||||
logger.info(f"Rank {sp_rank} Sum of output1: {sum_output1.item()}")
|
||||
logger.info(f"Rank {sp_rank} Sum of output2: {sum_output2.item()}")
|
||||
|
||||
# Clean up
|
||||
logger.info("Cleaning up distributed environment")
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
logger.info("Test completed successfully")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_hunyuanvideo_distributed()
|
||||
@@ -1,7 +1,9 @@
|
||||
The reference videos in the `reference_videos` directory are used as part of an e2e test to ensure consistency in video generation quality across code changes. `test_inference_similarity.py` compares newly generated videos against these references using Structural Similarity Index (SSIM) metrics to detect any regressions in visual quality across code changes.
|
||||
|
||||
`reference_videos/FLASH_ATTN/` videos were generated on commit `66107fd5b8469fed25972feb632cd48887dac451`.
|
||||
`reference_videos/TORCH_SDPA/` videos were generated on commit `4ea008b8a16d7f5678a44b187ebdd7d9d0416ff1`.
|
||||
`reference_videos/FastHunyuan-diffusers/FLASH_ATTN/` videos were generated on commit `66107fd5b8469fed25972feb632cd48887dac451`.
|
||||
`reference_videos/FastHunyuan-diffusers/TORCH_SDPA/` videos were generated on commit `4ea008b8a16d7f5678a44b187ebdd7d9d0416ff1`.
|
||||
`reference_videos/Wan2.1-T2V-1.3B-Diffusers` videos were generated on commit `d085770a70988c7b26632a0c3123c24a57f7ca77`.
|
||||
`reference_videos/Wan2.1-I2V-14B-480P-Diffusers` videos were generated on commit `d085770a70988c7b26632a0c3123c24a57f7ca77`.
|
||||
|
||||
## Generation Details
|
||||
|
||||
@@ -9,7 +11,7 @@ The reference videos in the `reference_videos` directory are used as part of an
|
||||
|
||||
## Generation Parameters
|
||||
|
||||
{
|
||||
FastHunyuan-diffusers: {
|
||||
"num_gpus": 2,
|
||||
"model_path": "data/FastHunyuan-diffusers",
|
||||
"height": 720,
|
||||
@@ -26,8 +28,51 @@ The reference videos in the `reference_videos` directory are used as part of an
|
||||
"fps": 24
|
||||
}
|
||||
|
||||
### Prompts
|
||||
Wan2.1-T2V-1.3B-Diffusers: {
|
||||
"num_gpus": 2,
|
||||
"model_path": "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 20,
|
||||
"guidance_scale": 3,
|
||||
"embedded_cfg_scale": 6,
|
||||
"flow_shift": 7.0,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 2,
|
||||
"vae_sp": True,
|
||||
"fps": 24,
|
||||
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
"text-encoder-precision": "fp32"
|
||||
}
|
||||
|
||||
1. 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."
|
||||
Wan2.1-I2V-14B-480P-Diffusers: {
|
||||
"num_gpus": 2,
|
||||
"model_path": "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 45,
|
||||
"num_inference_steps": 6,
|
||||
"guidance_scale": 5.0,
|
||||
"embedded_cfg_scale": 6,
|
||||
"flow_shift": 7.0,
|
||||
"seed": 1024,
|
||||
"sp_size": 2,
|
||||
"tp_size": 2,
|
||||
"vae_sp": True,
|
||||
"fps": 24,
|
||||
"neg_prompt": "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
"text-encoder-precision": "fp32"
|
||||
}
|
||||
|
||||
### Text-to-Video Prompts
|
||||
|
||||
1. "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."
|
||||
|
||||
2. "A lone hiker stands atop a towering cliff, silhouetted against the vast horizon. The rugged landscape stretches endlessly beneath, its earthy tones blending into the soft blues of the sky. The scene captures the spirit of exploration and human resilience. High angle, dynamic framing, with soft natural lighting emphasizing the grandeur of nature."
|
||||
|
||||
### Image-to-Video Prompts
|
||||
|
||||
1. "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
|
||||
Image path: "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user