Compare commits
18
Commits
docs-build
..
dmd
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
eb16ffb7b8 | ||
|
|
acd65e42e4 | ||
|
|
892d13d14d | ||
|
|
3a93f954df | ||
|
|
3e5ff5b583 | ||
|
|
84a6c32c96 | ||
|
|
ec9dd96d2e | ||
|
|
ed74b3c65c | ||
|
|
fce56124d7 | ||
|
|
4109928d27 | ||
|
|
d4ca37df9e | ||
|
|
9099e88e9a | ||
|
|
b679c8e515 | ||
|
|
da6003bc50 | ||
|
|
e811130464 | ||
|
|
d8a45e71c1 | ||
|
|
7b50887e38 | ||
|
|
89add12d3c |
@@ -1 +0,0 @@
|
||||
blank_issues_enabled: false
|
||||
@@ -1,240 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
|
||||
import requests
|
||||
|
||||
|
||||
def parse_arguments():
|
||||
"""Parse command line arguments"""
|
||||
parser = argparse.ArgumentParser(description='Run tests on RunPod GPU')
|
||||
parser.add_argument('--gpu-type', type=str, help='GPU type to use')
|
||||
parser.add_argument('--gpu-count',
|
||||
type=int,
|
||||
help='Number of GPUs to use',
|
||||
default=1)
|
||||
parser.add_argument('--test-command', type=str, help='Test command to run')
|
||||
parser.add_argument('--disk-size',
|
||||
type=int,
|
||||
default=20,
|
||||
help='Container disk size in GB (default: 20)')
|
||||
parser.add_argument('--volume-size',
|
||||
type=int,
|
||||
default=20,
|
||||
help='Persistent volume size in GB (default: 20)')
|
||||
parser.add_argument(
|
||||
'--image',
|
||||
type=str,
|
||||
default='runpod/pytorch:2.4.0-py3.11-cuda12.4.1-devel-ubuntu22.04',
|
||||
help='Docker image to use')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
args = parse_arguments()
|
||||
API_KEY = os.environ['RUNPOD_API_KEY']
|
||||
RUN_ID = os.environ['GITHUB_RUN_ID']
|
||||
JOB_ID = os.environ['JOB_ID']
|
||||
PODS_API = "https://rest.runpod.io/v1/pods"
|
||||
HEADERS = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {API_KEY}"
|
||||
}
|
||||
|
||||
|
||||
def create_pod():
|
||||
"""Create a RunPod instance"""
|
||||
print(f"Creating RunPod instance with GPU: {args.gpu_type}...")
|
||||
payload = {
|
||||
"name": f"fastvideo-{JOB_ID}-{RUN_ID}",
|
||||
"containerDiskInGb": args.disk_size,
|
||||
"volumeInGb": args.volume_size,
|
||||
"gpuTypeIds": [args.gpu_type],
|
||||
"gpuCount": args.gpu_count,
|
||||
"imageName": args.image,
|
||||
"allowedCudaVersions": ["12.4"]
|
||||
}
|
||||
|
||||
response = requests.post(PODS_API, headers=HEADERS, json=payload)
|
||||
response_data = response.json()
|
||||
print(f"Response: {json.dumps(response_data, indent=2)}")
|
||||
|
||||
return response_data["id"]
|
||||
|
||||
|
||||
def wait_for_pod(pod_id):
|
||||
"""Wait for pod to be in RUNNING state and fully ready with SSH access"""
|
||||
print("Waiting for RunPod to be ready...")
|
||||
|
||||
# First wait for RUNNING status
|
||||
max_attempts = 10
|
||||
attempts = 0
|
||||
while attempts < max_attempts:
|
||||
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
pod_data = response.json()
|
||||
status = pod_data["desiredStatus"]
|
||||
|
||||
if status == "RUNNING":
|
||||
print("RunPod is running! Now waiting for ports to be assigned...")
|
||||
break
|
||||
|
||||
print(
|
||||
f"Current status: {status}, waiting... (attempt {attempts+1}/{max_attempts})"
|
||||
)
|
||||
time.sleep(2)
|
||||
attempts += 1
|
||||
|
||||
if attempts >= max_attempts:
|
||||
raise TimeoutError(
|
||||
"Timed out waiting for RunPod to reach RUNNING state")
|
||||
|
||||
# Wait for ports to be assigned
|
||||
max_attempts = 6
|
||||
attempts = 0
|
||||
while attempts < max_attempts:
|
||||
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
pod_data = response.json()
|
||||
port_mappings = pod_data.get("portMappings")
|
||||
|
||||
if (port_mappings is not None and "22" in port_mappings
|
||||
and pod_data.get("publicIp", "") != ""):
|
||||
print("RunPod is ready with SSH access!")
|
||||
print(f"SSH IP: {pod_data['publicIp']}")
|
||||
print(f"SSH Port: {port_mappings['22']}")
|
||||
break
|
||||
|
||||
print(
|
||||
f"Waiting for SSH port and public IP to be available... (attempt {attempts+1}/{max_attempts})"
|
||||
)
|
||||
time.sleep(10)
|
||||
attempts += 1
|
||||
|
||||
if attempts >= max_attempts:
|
||||
raise TimeoutError("Timed out waiting for RunPod SSH access")
|
||||
|
||||
|
||||
def execute_command(pod_id):
|
||||
"""Execute command on the pod via SSH using system SSH client"""
|
||||
print(f"Running command: {args.test_command}")
|
||||
|
||||
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
pod_data = response.json()
|
||||
ssh_ip = pod_data["publicIp"]
|
||||
ssh_port = pod_data["portMappings"]["22"]
|
||||
|
||||
# Copy the repository to the pod using scp
|
||||
repo_dir = os.path.abspath(os.getcwd())
|
||||
repo_name = os.path.basename(repo_dir)
|
||||
|
||||
print(f"Copying repository from {repo_dir} to RunPod...")
|
||||
|
||||
tar_command = [
|
||||
"tar", "-czf", "/tmp/repo.tar.gz", "-C",
|
||||
os.path.dirname(repo_dir), repo_name
|
||||
]
|
||||
subprocess.run(tar_command, check=True)
|
||||
|
||||
# Copy the tarball to the pod
|
||||
scp_command = [
|
||||
"scp", "-o", "StrictHostKeyChecking=no", "-o",
|
||||
"UserKnownHostsFile=/dev/null", "-o", "ServerAliveInterval=60", "-o",
|
||||
"ServerAliveCountMax=10", "-P",
|
||||
str(ssh_port), "/tmp/repo.tar.gz", f"root@{ssh_ip}:/tmp/"
|
||||
]
|
||||
subprocess.run(scp_command, check=True)
|
||||
|
||||
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
|
||||
]
|
||||
remote_command = " && ".join(setup_steps)
|
||||
|
||||
ssh_command = [
|
||||
"ssh", "-o", "StrictHostKeyChecking=no", "-o",
|
||||
"UserKnownHostsFile=/dev/null", "-o", "ServerAliveInterval=60", "-o",
|
||||
"ServerAliveCountMax=10", "-p",
|
||||
str(ssh_port), f"root@{ssh_ip}", remote_command
|
||||
]
|
||||
|
||||
print(f"Connecting to {ssh_ip}:{ssh_port}...")
|
||||
|
||||
try:
|
||||
process = subprocess.Popen(ssh_command,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=0)
|
||||
|
||||
stdout_lines = []
|
||||
|
||||
print("Command output:")
|
||||
|
||||
for line in iter(process.stdout.readline, ''):
|
||||
print(line.strip())
|
||||
stdout_lines.append(line)
|
||||
|
||||
process.wait()
|
||||
|
||||
return_code = process.returncode
|
||||
success = return_code == 0
|
||||
|
||||
stdout_str = "".join(stdout_lines)
|
||||
|
||||
if success:
|
||||
print("Command executed successfully")
|
||||
else:
|
||||
print(f"Command failed with exit code {return_code}")
|
||||
|
||||
result = {
|
||||
"success": success,
|
||||
"return_code": return_code,
|
||||
"stdout": stdout_str,
|
||||
"stderr": ""
|
||||
}
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error executing SSH command: {str(e)}")
|
||||
result = {"success": False, "error": str(e), "stdout": "", "stderr": ""}
|
||||
return result
|
||||
|
||||
|
||||
def terminate_pod(pod_id):
|
||||
"""Terminate the pod"""
|
||||
print("Terminating RunPod...")
|
||||
requests.delete(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
print(f"Terminated pod {pod_id}")
|
||||
|
||||
|
||||
def main():
|
||||
pod_id = None
|
||||
try:
|
||||
pod_id = create_pod()
|
||||
wait_for_pod(pod_id)
|
||||
result = execute_command(pod_id)
|
||||
|
||||
if result.get("error") is not None:
|
||||
print(f"Error executing command: {result['error']}")
|
||||
sys.exit(1)
|
||||
|
||||
if not result.get("success", False):
|
||||
print(
|
||||
"Tests failed - check the output above for details on which tests failed"
|
||||
)
|
||||
sys.exit(1)
|
||||
|
||||
finally:
|
||||
if pod_id:
|
||||
terminate_pod(pod_id)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,90 +0,0 @@
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import uuid
|
||||
|
||||
import requests
|
||||
|
||||
API_KEY = os.environ['RUNPOD_API_KEY']
|
||||
RUN_ID = os.environ.get('GITHUB_RUN_ID', str(uuid.uuid4()))
|
||||
PODS_API = "https://rest.runpod.io/v1/pods"
|
||||
HEADERS = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {API_KEY}"
|
||||
}
|
||||
|
||||
|
||||
def get_job_ids():
|
||||
"""Parse job IDs from environment variable"""
|
||||
job_ids_str = os.environ.get('JOB_IDS')
|
||||
try:
|
||||
job_ids = json.loads(job_ids_str)
|
||||
if not isinstance(job_ids, list):
|
||||
print("Error: JOB_IDS is not a list.")
|
||||
sys.exit(1)
|
||||
return job_ids
|
||||
except json.JSONDecodeError as e:
|
||||
print(f"Error parsing JOB_IDS: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def cleanup_pods():
|
||||
"""Find and terminate RunPod instances"""
|
||||
print(f"Run ID: {RUN_ID}")
|
||||
|
||||
single_job_id = os.environ.get('JOB_ID')
|
||||
|
||||
if single_job_id:
|
||||
job_ids = [single_job_id]
|
||||
print(f"Job ID: {single_job_id}")
|
||||
else:
|
||||
job_ids = get_job_ids()
|
||||
print(f"Job IDs: {job_ids}")
|
||||
|
||||
# Get all pods associated with RunPod API_KEY
|
||||
try:
|
||||
response = requests.get(PODS_API, headers=HEADERS)
|
||||
response.raise_for_status()
|
||||
pods = response.json()
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"Error getting pods: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
# Find and terminate pods created by this workflow run
|
||||
terminated_pods = []
|
||||
for pod in pods:
|
||||
pod_name = pod.get("name", "")
|
||||
pod_id = pod.get("id")
|
||||
|
||||
# Check if this pod was created by one of our jobs
|
||||
if any(f"{job_id}-{RUN_ID}" in pod_name for job_id in job_ids):
|
||||
print(f"Found pod: {pod_id} ({pod_name})")
|
||||
try:
|
||||
print(f"Terminating pod {pod_id}...")
|
||||
term_response = requests.delete(f"{PODS_API}/{pod_id}",
|
||||
headers=HEADERS)
|
||||
term_response.raise_for_status()
|
||||
terminated_pods.append(pod_id)
|
||||
print(f"Successfully terminated pod {pod_id}")
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"Error terminating pod {pod_id}: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
if terminated_pods:
|
||||
if single_job_id:
|
||||
print(f"Terminated pod: {terminated_pods[0]}")
|
||||
else:
|
||||
print(f"Terminated {len(terminated_pods)} pods: {terminated_pods}")
|
||||
else:
|
||||
if single_job_id:
|
||||
print(f"No pod found matching pattern: {single_job_id}-{RUN_ID}")
|
||||
else:
|
||||
print("No pods found to terminate.")
|
||||
|
||||
|
||||
def main():
|
||||
cleanup_pods()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,77 +0,0 @@
|
||||
# Sample workflow for building and deploying a Hugo site to GitHub Pages
|
||||
name: Deploy FastVideo Docs to Pages
|
||||
|
||||
on:
|
||||
# Runs on pushes targeting the default branch
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "docs/**/*.md"
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
types: [opened, ready_for_review, synchronize, reopened]
|
||||
paths:
|
||||
- "docs/**/*.md"
|
||||
|
||||
# Allows you to run this workflow manually from the Actions tab
|
||||
workflow_dispatch:
|
||||
|
||||
# Sets permissions of the GITHUB_TOKEN to allow deployment to GitHub Pages
|
||||
permissions:
|
||||
contents: read
|
||||
pages: write
|
||||
id-token: write
|
||||
|
||||
# Allow only one concurrent deployment, skipping runs queued between the run in-progress and latest queued.
|
||||
# However, do NOT cancel in-progress runs as we want to allow these production deployments to complete.
|
||||
concurrency:
|
||||
group: "pages"
|
||||
cancel-in-progress: false
|
||||
|
||||
# Default to bash
|
||||
defaults:
|
||||
run:
|
||||
shell: bash
|
||||
|
||||
jobs:
|
||||
# Build job
|
||||
build:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
- name: Setup Pages
|
||||
id: pages
|
||||
uses: actions/configure-pages@v5
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
cd docs
|
||||
pip install -r requirements-docs.txt
|
||||
- name: Build docs
|
||||
run: |
|
||||
cd docs
|
||||
make clean
|
||||
make html
|
||||
- name: Upload artifact
|
||||
uses: actions/upload-pages-artifact@v3
|
||||
with:
|
||||
path: ./docs/build/html
|
||||
|
||||
# Deployment job
|
||||
deploy:
|
||||
environment:
|
||||
name: github-pages
|
||||
url: ${{ steps.deployment.outputs.page_url }}
|
||||
if: ${{ github.event_name == 'push' }}
|
||||
runs-on: ubuntu-latest
|
||||
needs: build
|
||||
steps:
|
||||
- name: Deploy to GitHub Pages
|
||||
id: deployment
|
||||
uses: actions/deploy-pages@v4
|
||||
@@ -1,70 +0,0 @@
|
||||
name: Publish FastVideo to PyPI on Version Change
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- 'pyproject.toml' # Trigger when pyproject.toml changes
|
||||
|
||||
jobs:
|
||||
check-version-change:
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
version-changed: ${{ steps.check-version.outputs.changed }}
|
||||
new-version: ${{ steps.check-version.outputs.new-version }}
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 2
|
||||
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP "version\\s*=\\s*\"\\K[^\"]+\"" pyproject.toml)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./pyproject.toml | grep -oP "version\\s*=\\s*\"\\K[^\"]+\"" || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
|
||||
echo "changed=true" >> $GITHUB_OUTPUT
|
||||
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
|
||||
else
|
||||
echo "Version did not change"
|
||||
echo "changed=false" >> $GITHUB_OUTPUT
|
||||
fi
|
||||
|
||||
build-publish-main:
|
||||
needs: check-version-change
|
||||
if: needs.check-version-change.outputs.version-changed == 'true'
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
- name: Install build dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install build twine wheel
|
||||
|
||||
- name: Build package
|
||||
run: |
|
||||
python -m build
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: dist/
|
||||
@@ -1,17 +0,0 @@
|
||||
{
|
||||
"problemMatcher": [
|
||||
{
|
||||
"owner": "actionlint",
|
||||
"pattern": [
|
||||
{
|
||||
"regexp": "^(?:\\x1b\\[\\d+m)?(.+?)(?:\\x1b\\[\\d+m)*:(?:\\x1b\\[\\d+m)*(\\d+)(?:\\x1b\\[\\d+m)*:(?:\\x1b\\[\\d+m)*(\\d+)(?:\\x1b\\[\\d+m)*: (?:\\x1b\\[\\d+m)*(.+?)(?:\\x1b\\[\\d+m)* \\[(.+?)\\]$",
|
||||
"file": 1,
|
||||
"line": 2,
|
||||
"column": 3,
|
||||
"message": 4,
|
||||
"code": 5
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -1,16 +0,0 @@
|
||||
{
|
||||
"problemMatcher": [
|
||||
{
|
||||
"owner": "mypy",
|
||||
"pattern": [
|
||||
{
|
||||
"regexp": "^(.+):(\\d+):\\s(error|warning):\\s(.+)$",
|
||||
"file": 1,
|
||||
"line": 2,
|
||||
"severity": 3,
|
||||
"message": 4
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -1,173 +0,0 @@
|
||||
name: PR Test
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: [main]
|
||||
paths:
|
||||
- "fastvideo/**/*.py"
|
||||
- ".github/workflows/pr-test.yml"
|
||||
pull_request:
|
||||
branches: [main]
|
||||
types: [opened, ready_for_review, synchronize, reopened]
|
||||
paths:
|
||||
- "fastvideo/**/*.py"
|
||||
- ".github/workflows/pr-test.yml"
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
run_encoder_test:
|
||||
description: "Run encoder-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_ssim_test:
|
||||
description: "Run ssim-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
|
||||
concurrency:
|
||||
group: pr-test-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
uses: ./.github/workflows/pre-commit.yml
|
||||
|
||||
change-filter:
|
||||
runs-on: ubuntu-latest
|
||||
needs: pre-commit
|
||||
if: ${{ github.event.pull_request.draft == false || github.event_name == 'workflow_dispatch' }}
|
||||
outputs:
|
||||
encoder-test: ${{ steps.filter.outputs.encoder-test }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
id: filter
|
||||
with:
|
||||
filters: |
|
||||
encoder-test:
|
||||
- 'fastvideo/v1/models/encoders/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/encoders/**'
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.encoder-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_encoder_test == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
environment: runpod-runners
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
- name: Set up SSH key
|
||||
run: |
|
||||
mkdir -p ~/.ssh
|
||||
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
|
||||
chmod 600 ~/.ssh/id_rsa
|
||||
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
|
||||
|
||||
- name: Install dependencies
|
||||
run: pip install requests
|
||||
|
||||
- name: Run tests on RunPod
|
||||
env:
|
||||
JOB_ID: "encoder-test"
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
timeout-minutes: 30
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
--gpu-type "NVIDIA A40"
|
||||
--gpu-count 1
|
||||
--volume-size 100
|
||||
--test-command "pip install -e .[test] &&
|
||||
pip install flash-attn==2.7.0.post2 --no-build-isolation &&
|
||||
pytest ./fastvideo/v1/tests/encoders -s"
|
||||
|
||||
- name: Terminate RunPod Instances
|
||||
if: ${{ always() }}
|
||||
env:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
JOB_ID: "encoder-test"
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
ssim-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
environment: runpod-runners
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
- name: Set up SSH key
|
||||
run: |
|
||||
mkdir -p ~/.ssh
|
||||
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
|
||||
chmod 600 ~/.ssh/id_rsa
|
||||
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
|
||||
|
||||
- name: Install dependencies
|
||||
run: pip install requests
|
||||
|
||||
- name: Run tests on RunPod
|
||||
env:
|
||||
JOB_ID: "ssim-test"
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
timeout-minutes: 30
|
||||
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"
|
||||
|
||||
- name: Terminate RunPod Instances
|
||||
if: ${{ always() }}
|
||||
env:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
JOB_ID: "ssim-test"
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
runpod-cleanup:
|
||||
needs: [encoder-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:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
- name: Install dependencies
|
||||
run: pip install requests
|
||||
|
||||
- name: Cleanup all RunPod instances
|
||||
env:
|
||||
JOB_IDS: '["encoder-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
|
||||
@@ -1,18 +0,0 @@
|
||||
name: pre-commit
|
||||
|
||||
on:
|
||||
workflow_call:
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/actionlint.json"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/mypy.json"
|
||||
- uses: pre-commit/action@v3.0.1
|
||||
with:
|
||||
extra_args: --all-files --hook-stage manual
|
||||
@@ -1,221 +0,0 @@
|
||||
name: Publish Sliding Tile Attention Kernel to PyPI on Version Change
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/sliding_tile_attention/setup.py"
|
||||
|
||||
jobs:
|
||||
check-version-change:
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
version-changed: ${{ steps.check-version.outputs.changed }}
|
||||
new-version: ${{ steps.check-version.outputs.new-version }}
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 2
|
||||
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/sliding_tile_attention
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
|
||||
echo "changed=true" >> $GITHUB_OUTPUT
|
||||
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
|
||||
else
|
||||
echo "Version did not change"
|
||||
echo "changed=false" >> $GITHUB_OUTPUT
|
||||
fi
|
||||
|
||||
build_wheels:
|
||||
name: Build Wheel
|
||||
needs: check-version-change
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' }}
|
||||
runs-on: ${{ matrix.os }}
|
||||
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
# Using ubuntu-20.04 instead of 22.04 for more compatibility (glibc). Ideally we'd use the
|
||||
# manylinux docker image, but I haven't figured out how to install CUDA on manylinux.
|
||||
os: [ubuntu-22.04]
|
||||
python-version: ['3.10', '3.11', '3.12', '3.13']
|
||||
torch-version: ['2.5.1', '2.6.0']
|
||||
cuda-version: ['12.4.1', '12.5.1', '12.6.3']
|
||||
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Install CUDA ${{ matrix.cuda-version }}
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: ${{ matrix.cuda-version }}
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-${{ matrix.cuda-version }}
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch ${{ matrix.torch-version }}+cu${{ matrix.cuda-version }}
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
|
||||
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
|
||||
pip install typing-extensions==4.12.2
|
||||
# We want to figure out the CUDA version to download pytorch
|
||||
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
|
||||
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
|
||||
export TORCH_CUDA_VERSION=124
|
||||
pip install --no-cache-dir torch==${{ matrix.torch-version }} --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build wheel
|
||||
run: |
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/sliding_tile_attention # Move into the correct folder
|
||||
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/sliding_tile_attention
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-version }} | cut -d. -f1,2)
|
||||
# Get the correct version format
|
||||
tmpname=cu${CUDA_SHORT_VERSION}torch${TORCH_SHORT_VERSION}
|
||||
wheel_name=$(ls dist/*whl | xargs -n 1 basename | sed "s/-/+$tmpname-/2")
|
||||
# Rename with version information
|
||||
ls dist/*whl |xargs -I {} mv {} dist/${wheel_name}
|
||||
echo "wheel_name=${wheel_name}" >> $GITHUB_ENV
|
||||
|
||||
- name: Upload wheel artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}
|
||||
path: csrc/sliding_tile_attention/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
name: Publish package
|
||||
needs: [build_wheels, check-version-change]
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' }}
|
||||
runs-on: ubuntu-22.04
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
- name: Install CUDA 12.4.1
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: 12.4.1
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
sub-packages: '["nvcc"]'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-12.4.1
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch 2.5.1+cu12.4.1
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
|
||||
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
|
||||
pip install typing-extensions==4.12.2
|
||||
# We want to figure out the CUDA version to download pytorch
|
||||
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
|
||||
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
|
||||
export TORCH_CUDA_VERSION=124
|
||||
pip install --no-cache-dir torch==2.5.1 --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/sliding_tile_attention # Move into the correct folder
|
||||
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/sliding_tile_attention/dist/
|
||||
@@ -1,31 +0,0 @@
|
||||
name: Run Tests
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: [ main ]
|
||||
pull_request:
|
||||
branches: [ main ]
|
||||
|
||||
jobs:
|
||||
test:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.12' # or any version you need
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip setuptools wheel
|
||||
pip install torch
|
||||
pip install packaging ninja
|
||||
pip install -e .
|
||||
pip install pytest
|
||||
|
||||
- name: Run Pytest
|
||||
run: |
|
||||
pytest --ignore csrc/sliding_tile_attention/test
|
||||
+27
-9
@@ -1,4 +1,6 @@
|
||||
ucf101_stride4x4x4
|
||||
__pycache__
|
||||
*.mp4
|
||||
.ipynb_checkpoints
|
||||
*.pth
|
||||
UCF-101/
|
||||
@@ -6,17 +8,41 @@ results/
|
||||
build/
|
||||
fastvideo.egg-info/
|
||||
wandb/
|
||||
.idea
|
||||
*.ipynb
|
||||
*.jpg
|
||||
*.mp3
|
||||
*.safetensors
|
||||
*.mp4
|
||||
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
|
||||
*.png
|
||||
*.gif
|
||||
*.pth
|
||||
*.pt
|
||||
cache_dir/
|
||||
wandb/
|
||||
test*
|
||||
sample_video*
|
||||
sample_image*
|
||||
512*
|
||||
720*
|
||||
1024*
|
||||
debug*
|
||||
private*
|
||||
caption*
|
||||
*deepspeed*
|
||||
revised*
|
||||
129f*
|
||||
all*
|
||||
read*
|
||||
YSH*
|
||||
*pick*
|
||||
*ysh*
|
||||
hw*
|
||||
257f*
|
||||
513f*
|
||||
taming*
|
||||
221hw*
|
||||
65x512x512
|
||||
runs/
|
||||
samples/
|
||||
*validation/
|
||||
@@ -26,11 +52,3 @@ outputs_video
|
||||
sbatch.sh
|
||||
*.out
|
||||
env
|
||||
dist/
|
||||
*.o
|
||||
**/build/
|
||||
**.egg-info
|
||||
**.pyc
|
||||
**.egg
|
||||
**.txt
|
||||
**.json
|
||||
@@ -1,3 +0,0 @@
|
||||
[submodule "csrc/sliding_tile_attention/tk"]
|
||||
path = csrc/sliding_tile_attention/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
@@ -1,80 +0,0 @@
|
||||
default_stages:
|
||||
- pre-commit # Run locally
|
||||
- manual # Run in CI
|
||||
exclude: |
|
||||
(?x)(
|
||||
fastvideo/v1/third_party/.*|
|
||||
csrc/.*|
|
||||
assets/.*|
|
||||
tests/.*|
|
||||
demo/.*|
|
||||
predict\.py|
|
||||
scripts/.*|
|
||||
fastvideo/data_preprocess/.*|
|
||||
fastvideo/dataset/.*|
|
||||
fastvideo/distill/.*|
|
||||
fastvideo/distill\.py|
|
||||
fastvideo/distill_adv\.py|
|
||||
fastvideo/models/.*|
|
||||
fastvideo/sample/.*|
|
||||
fastvideo/train\.py|
|
||||
fastvideo/utils/.*|
|
||||
.github/workflows/fastvideo-publish.yml|
|
||||
.github/workflows/sta-publish.yml
|
||||
)
|
||||
repos:
|
||||
- repo: https://github.com/google/yapf
|
||||
rev: v0.43.0
|
||||
hooks:
|
||||
- id: yapf
|
||||
args: [--in-place, --verbose]
|
||||
additional_dependencies: [toml] # TODO: Remove when yapf is upgraded
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
rev: v0.11.4
|
||||
hooks:
|
||||
- id: ruff
|
||||
args: [--output-format, github, --fix]
|
||||
- repo: https://github.com/codespell-project/codespell
|
||||
rev: v2.4.1
|
||||
hooks:
|
||||
- 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/jackdewinter/pymarkdown
|
||||
rev: v0.9.29
|
||||
hooks:
|
||||
- id: pymarkdown
|
||||
args: [fix]
|
||||
- repo: https://github.com/rhysd/actionlint
|
||||
rev: v1.7.7
|
||||
hooks:
|
||||
- id: actionlint
|
||||
- repo: https://github.com/pre-commit/mirrors-mypy
|
||||
rev: v1.15.0
|
||||
hooks:
|
||||
- id: mypy
|
||||
args: [--python-version, '3.10', --follow-imports, "skip", ]
|
||||
additional_dependencies: [types-cachetools, types-setuptools, types-PyYAML, types-requests]
|
||||
- repo: local
|
||||
hooks:
|
||||
- id: check-filenames
|
||||
name: Check for spaces in all filenames
|
||||
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'
|
||||
language: system
|
||||
always_run: true
|
||||
pass_filenames: false
|
||||
# Keep `suggestion` last
|
||||
- id: suggestion
|
||||
name: Suggestion
|
||||
entry: bash -c 'echo "To bypass pre-commit hooks, add --no-verify to git commit."'
|
||||
language: system
|
||||
verbose: true
|
||||
pass_filenames: false
|
||||
# Insert new entries above the `suggestion` entry
|
||||
@@ -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
|
||||
@@ -184,4 +184,18 @@
|
||||
comment syntax for the file format. We also recommend that a
|
||||
file or class name and description of purpose be included on the
|
||||
same "printed page" as the copyright notice for easier
|
||||
identification within third-party archives.
|
||||
identification within third-party archives.
|
||||
|
||||
Copyright [2023] Lightning AI
|
||||
|
||||
Licensed under the Apache License, Version 2.0 (the "License");
|
||||
you may not use this file except in compliance with the License.
|
||||
You may obtain a copy of the License at
|
||||
|
||||
http://www.apache.org/licenses/LICENSE-2.0
|
||||
|
||||
Unless required by applicable law or agreed to in writing, software
|
||||
distributed under the License is distributed on an "AS IS" BASIS,
|
||||
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
See the License for the specific language governing permissions and
|
||||
limitations under the License.
|
||||
@@ -4,15 +4,19 @@
|
||||
|
||||
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>
|
||||
</p>
|
||||
|
||||
https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1
|
||||
https://github.com/user-attachments/assets/5fbc4596-56d6-43aa-98e0-da472cf8e26c
|
||||
|
||||
|
||||
|
||||
|
||||
<p align="center">
|
||||
🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank">FastMochi</a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank">FastHunyuan</a> | 🎮 <a href="https://discord.gg/REBzDQTWWt" target="_blank"> Discord </a> | 🕹️ <a href="https://replicate.com/lucataco/fast-hunyuan-video" target="_blank"> Replicate </a>
|
||||
</p>
|
||||
|
||||
|
||||
FastVideo currently offers: (with more to come)
|
||||
|
||||
- [NEW!] [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- FastHunyuan and FastMochi: consistency distilled video diffusion models for 8x inference speedup.
|
||||
- First open distillation recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
|
||||
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
|
||||
@@ -21,89 +25,61 @@ FastVideo currently offers: (with more to come)
|
||||
|
||||
Dev in progress and highly experimental.
|
||||
|
||||
## 🎥 More Demos
|
||||
|
||||
Fast-Hunyuan comparison with original Hunyuan, achieving an 8X diffusion speed boost with the FastVideo framework.
|
||||
|
||||
https://github.com/user-attachments/assets/064ac1d2-11ed-4a0c-955b-4d412a96ef30
|
||||
|
||||
Comparison between OpenAI Sora, original Hunyuan and FastHunyuan
|
||||
|
||||
https://github.com/user-attachments/assets/d323b712-3f68-42b2-952b-94f6a49c4836
|
||||
|
||||
Comparison between original FastHunyuan, LLM-INT8 quantized FastHunyuan and NF4 quantized FastHunyuan
|
||||
|
||||
https://github.com/user-attachments/assets/cf89efb5-5f68-4949-a085-f41c1ef26c94
|
||||
|
||||
## Change Log
|
||||
- ```2025/02/20```: FastVideo now supports STA on [StepVideo](https://github.com/stepfun-ai/Step-Video-T2V) with 3.4X speedup!
|
||||
- ```2025/02/18```: Release the inference code and kernel for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- ```2025/01/13```: Support Lora finetuning for HunyuanVideo.
|
||||
- ```2024/12/25```: Enable single 4090 inference for `FastHunyuan`, please rerun the installation steps to update the environment.
|
||||
- ```2024/12/17```: `FastVideo` v1.0 is released.
|
||||
|
||||
## 🔧 Installation from source
|
||||
The code is tested on Python 3.10.0, CUDA 12.4 and H100.
|
||||
|
||||
## 🔧 Installation
|
||||
The code is tested on Python 3.10.0, CUDA 12.1 and H100.
|
||||
```
|
||||
# Clone FastVideo
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
|
||||
# Install FastVideo
|
||||
pip install -e .
|
||||
|
||||
# Install Flash Attention (optional)
|
||||
pip install flash-attn==2.7.0.post2
|
||||
./env_setup.sh fastvideo
|
||||
```
|
||||
|
||||
To try Sliding Tile Attention (optional), please follow the instruction in [csrc/sliding_tile_attention/README.md](csrc/sliding_tile_attention/README.md) to install STA.
|
||||
|
||||
## 🚀 Inference
|
||||
### Inference StepVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
```
|
||||
python scripts/huggingface/download_hf.py --repo_id=stepfun-ai/stepvideo-t2v --local_dir=data/stepvideo-t2v --repo_type=model
|
||||
```
|
||||
|
||||
Use the following scripts to run inference for StepVideo. When using STA for inference, the generated videos will have dimensions of 204×768×768 (currently, this is the only supported shape).
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_stepvideo_STA.sh # Inference stepvideo with STA
|
||||
sh scripts/inference/inference_stepvideo.sh # Inference original stepvideo
|
||||
```
|
||||
|
||||
### Inference HunyuanVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model
|
||||
```
|
||||
|
||||
We provide two examples in the following script to run inference with STA + [TeaCache](https://github.com/ali-vilab/TeaCache) and STA only.
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_hunyuan_STA.sh
|
||||
```
|
||||
|
||||
### Video Demos using STA + Teacache
|
||||
Visit our [demo website](https://fast-video.github.io/) to explore our complete collection of examples. We shorten a single video generation process from 945s to 317s on H100.
|
||||
|
||||
### Inference FastHunyuan on single RTX4090
|
||||
We now support NF4 and LLM-INT8 quantized inference using BitsAndBytes for FastHunyuan. With NF4 quantization, inference can be performed on a single RTX 4090 GPU, requiring just 20GB of VRAM.
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan-diffusers --local_dir=data/FastHunyuan-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan_hf_quantization.sh
|
||||
bash scripts/inference/inference_diffusers_hunyuan.sh
|
||||
```
|
||||
|
||||
For more information about the VRAM requirements for BitsAndBytes quantization, please refer to the table below (timing measured on an H100 GPU):
|
||||
|
||||
|
||||
| Configuration | Memory to Init Transformer | Peak Memory After Init Pipeline (Denoise) | Diffusion Time | End-to-End Time |
|
||||
|--------------------------------|----------------------------|--------------------------------------------|----------------|-----------------|
|
||||
| BF16 + Pipeline CPU Offload | 23.883G | 33.744G | 81s | 121.5s |
|
||||
| INT8 + Pipeline CPU Offload | 13.911G | 27.979G | 88s | 116.7s |
|
||||
| NF4 + Pipeline CPU Offload | 9.453G | 19.26G | 78s | 114.5s |
|
||||
|
||||
|
||||
|
||||
For improved quality in generated videos, we recommend using a GPU with 80GB of memory to run the BF16 model with the original Hunyuan pipeline. To execute the inference, use the following section:
|
||||
|
||||
### FastHunyuan
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan --local_dir=data/FastHunyuan --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan.sh
|
||||
```
|
||||
|
||||
You can also inference FastHunyuan in the [official Hunyuan github](https://github.com/Tencent/HunyuanVideo).
|
||||
|
||||
### FastMochi
|
||||
@@ -115,100 +91,53 @@ python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastMochi-diffuser
|
||||
bash scripts/inference/inference_mochi_sp.sh
|
||||
```
|
||||
|
||||
|
||||
## 🎯 Distill
|
||||
Our distillation recipe is based on [Phased Consistency Model](https://github.com/G-U-N/Phased-Consistency-Model). We did not find significant improvement using multi-phase distillation, so we keep the one phase setup similar to the original latent consistency model's recipe.
|
||||
We use the [MixKit](https://huggingface.co/datasets/LanguageBind/Open-Sora-Plan-v1.1.0/tree/main/all_mixkit) dataset for distillation. To avoid running the text encoder and VAE during training, we preprocess all data to generate text embeddings and VAE latents.
|
||||
Preprocessing instructions can be found [data_preprocess.md](docs/data_preprocess.md). For convenience, we also provide preprocessed data that can be downloaded directly using the following command:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/HD-Mixkit-Finetune-Hunyuan --local_dir=data/HD-Mixkit-Finetune-Hunyuan --repo_type=dataset
|
||||
```
|
||||
|
||||
Next, download the original model weights with:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model # original hunyuan
|
||||
python scripts/huggingface/download_hf.py --repo_id=genmo/mochi-1-preview --local_dir=data/mochi --repo_type=model # original mochi
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model # original hunyuan
|
||||
```
|
||||
|
||||
To launch the distillation process, use the following commands:
|
||||
|
||||
```
|
||||
bash scripts/distill/distill_hunyuan.sh # for hunyuan
|
||||
bash scripts/distill/distill_mochi.sh # for mochi
|
||||
bash scripts/distill/distill_hunyuan.sh # for hunyuan
|
||||
```
|
||||
|
||||
We also provide an optional script for distillation with adversarial loss, located at `fastvideo/distill_adv.py`. Although we tried adversarial loss, we did not observe significant improvements.
|
||||
## Finetune
|
||||
### ⚡ Full Finetune
|
||||
Ensure your data is prepared and preprocessed in the format specified in [data_preprocess.md](docs/data_preprocess.md). For convenience, we also provide a mochi preprocessed Black Myth Wukong data that can be downloaded directly:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Mochi-Black-Myth --local_dir=data/Mochi-Black-Myth --repo_type=dataset
|
||||
```
|
||||
|
||||
Download the original model weights as specified in [Distill Section](#-distill):
|
||||
Download the original model weights as specificed in [Distill Section](#-distill):
|
||||
|
||||
Then you can run the finetune with:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script**
|
||||
### ⚡ Lora Finetune
|
||||
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight --local_dir=data/Hunyuan-Black-Myth-Wukong-lora-weight --repo_type=model
|
||||
Currently, we only provide Lora Finetune for Mochi model, the command for Lora Finetune is
|
||||
```
|
||||
|
||||
#### Minimum Hardware Requirement
|
||||
- 40 GB GPU memory each for 2 GPUs with lora.
|
||||
bash scripts/finetune/finetune_mochi_lora.sh
|
||||
```
|
||||
### Minimum Hardware Requirement
|
||||
- 40 GB GPU memory each for 2 GPUs with lora
|
||||
- 30 GB GPU memory each for 2 GPUs with CPU offload and lora.
|
||||
|
||||
Currently, both Mochi and Hunyuan models support Lora finetuning through diffusers. To generate personalized videos from your own dataset, you'll need to follow three main steps: dataset preparation, finetuning, and inference.
|
||||
|
||||
#### Dataset Preparation
|
||||
We provide scripts to better help you get started to train on your own characters!
|
||||
You can run this to organize your dataset to get the videos2caption.json before preprocess. Specify your video folder and corresponding caption folder (caption files should be .txt files and have the same name with its video):
|
||||
|
||||
```
|
||||
python scripts/dataset_preparation/prepare_json_file.py --video_dir data/input_videos/ --prompt_dir data/captions/ --output_path data/output_folder/videos2caption.json --verbose
|
||||
```
|
||||
|
||||
Also, we provide script to resize your videos:
|
||||
|
||||
```
|
||||
python scripts/data_preprocess/resize_videos.py
|
||||
```
|
||||
|
||||
#### Finetuning
|
||||
After basic dataset preparation and preprocess, you can start to finetune your model using Lora:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_hunyuan_hf_lora.sh
|
||||
```
|
||||
|
||||
#### Inference
|
||||
For inference with Lora checkpoint, you can run the following scripts with additional parameter `--lora_checkpoint_dir`:
|
||||
|
||||
```
|
||||
bash scripts/inference/inference_hunyuan_hf.sh
|
||||
```
|
||||
|
||||
**We also provide scripts for Mochi in the same directory.**
|
||||
|
||||
#### Finetune with Both Image and Video
|
||||
Our codebase support finetuning with both image and video.
|
||||
|
||||
### Finetune with Both Image and Video
|
||||
Our codebase support finetuning with both image and video.
|
||||
```bash
|
||||
bash scripts/finetune/finetune_hunyuan.sh
|
||||
bash scripts/finetune/finetune_mochi_lora_mix.sh
|
||||
```
|
||||
|
||||
For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` option in your script.
|
||||
For Image-Video Mixture Fine-tuning, make sure to enable the --group_frame option in your script.
|
||||
|
||||
## 📑 Development Plan
|
||||
|
||||
@@ -220,38 +149,7 @@ For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` opt
|
||||
- [ ] fp8 support
|
||||
- [ ] faster load model and save model support
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please run `bash format.sh --all` before submitting a pull request.
|
||||
|
||||
## 🔧 Testing
|
||||
Run `pytest` to verify the data preprocessing, checkpoint saving, and sequence parallel pipelines. We recommend adding corresponding test cases in the `test` folder to support your contribution.
|
||||
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects: [PCM](https://github.com/G-U-N/Phased-Consistency-Model), [diffusers](https://github.com/huggingface/diffusers), [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan), and [xDiT](https://github.com/xdit-project/xDiT).
|
||||
|
||||
We thank MBZUAI and Anyscale for their support throughout this project.
|
||||
|
||||
## Citation
|
||||
If you use FastVideo for your research, please cite our paper:
|
||||
|
||||
```bibtex
|
||||
@misc{zhang2025fastvideogenerationsliding,
|
||||
title={Fast Video Generation with Sliding Tile Attention},
|
||||
author={Peiyuan Zhang and Yongqi Chen and Runlong Su and Hangliang Ding and Ion Stoica and Zhenghong Liu and Hao Zhang},
|
||||
year={2025},
|
||||
eprint={2502.04507},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV},
|
||||
url={https://arxiv.org/abs/2502.04507},
|
||||
}
|
||||
@misc{ding2025efficientvditefficientvideodiffusion,
|
||||
title={Efficient-vDiT: Efficient Video Diffusion Transformers With Attention Tile},
|
||||
author={Hangliang Ding and Dacheng Li and Runlong Su and Peiyuan Zhang and Zhijie Deng and Ion Stoica and Hao Zhang},
|
||||
year={2025},
|
||||
eprint={2502.06155},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV},
|
||||
url={https://arxiv.org/abs/2502.06155},
|
||||
}
|
||||
```
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
Binary file not shown.
|
Before Width: | Height: | Size: 751 KiB |
@@ -1,2 +0,0 @@
|
||||
recursive-include tk *
|
||||
include config.py
|
||||
@@ -1,68 +0,0 @@
|
||||
|
||||
|
||||
# Sliding Tile Atteniton Kernel
|
||||
|
||||
|
||||
## Installation
|
||||
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
|
||||
First, install C++20 for ThunderKittens:
|
||||
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
Install STA:
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
|
||||
```
|
||||
|
||||
|
||||
## Test
|
||||
```bash
|
||||
python test/test_sta.py
|
||||
```
|
||||
|
||||
## How Does STA Work?
|
||||
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
|
||||
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
|
||||
|
||||
## Why is STA Fast?
|
||||
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
|
||||
|
||||
STA removes mixed blocks.
|
||||
|
||||
|
||||
<div align="center">
|
||||
<img src=../../assets/sliding_tile_attn_map.png width="80%"/>
|
||||
</div>
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
|
||||
@@ -1,15 +0,0 @@
|
||||
### ADD TO THIS TO REGISTER NEW KERNELS
|
||||
sources = {
|
||||
'attn': {
|
||||
'source_files': {
|
||||
'h100': 'st_attn/st_attn_h100.cu' # define these source files for each GPU target desired.
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
### WHICH KERNELS DO WE WANT TO BUILD?
|
||||
# (oftentimes during development work you don't need to redefine them all.)
|
||||
kernels = ['attn']
|
||||
|
||||
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
|
||||
target = 'h100'
|
||||
@@ -1,76 +0,0 @@
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
from config import kernels, sources, target
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "st_attn"
|
||||
VERSION = "0.0.2"
|
||||
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"
|
||||
|
||||
# Set environment variables
|
||||
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
|
||||
python_include = subprocess.check_output(['python', '-c',
|
||||
"import sysconfig; print(sysconfig.get_path('include'))"]).decode().strip()
|
||||
torch_include = subprocess.check_output([
|
||||
'python', '-c',
|
||||
"import torch; from torch.utils.cpp_extension import include_paths; print(' '.join(['-I' + p for p in include_paths()]))"
|
||||
]).decode().strip()
|
||||
print('st_attn root:', tk_root)
|
||||
print('Python include:', python_include)
|
||||
print('Torch include directories:', torch_include)
|
||||
|
||||
# CUDA flags
|
||||
cuda_flags = [
|
||||
'-DNDEBUG', '-Xcompiler=-Wno-psabi', '-Xcompiler=-fno-strict-aliasing', '--expt-extended-lambda',
|
||||
'--expt-relaxed-constexpr', '-forward-unknown-to-host-compiler', '--use_fast_math', '-std=c++20', '-O3',
|
||||
'-Xnvlink=--verbose', '-Xptxas=--verbose', '-Xptxas=--warn-on-spills', f'-I{tk_root}/include',
|
||||
f'-I{tk_root}/prototype', f'-I{python_include}', '-DTORCH_COMPILE'
|
||||
] + torch_include.split()
|
||||
cpp_flags = ['-std=c++20', '-O3']
|
||||
|
||||
if target == 'h100':
|
||||
cuda_flags.append('-DKITTENS_HOPPER')
|
||||
cuda_flags.append('-arch=sm_90a')
|
||||
else:
|
||||
raise ValueError(f'Target {target} not supported')
|
||||
|
||||
source_files = ['st_attn.cpp']
|
||||
for k in kernels:
|
||||
if target not in sources[k]['source_files']:
|
||||
raise KeyError(f'Target {target} not found in source files for kernel {k}')
|
||||
if isinstance(sources[k]['source_files'][target], list):
|
||||
source_files.extend(sources[k]['source_files'][target])
|
||||
else:
|
||||
source_files.append(sources[k]['source_files'][target])
|
||||
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
|
||||
|
||||
setup(name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
ext_modules=[
|
||||
CUDAExtension('st_attn_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
],
|
||||
cmdclass={'build_ext': BuildExtension},
|
||||
classifiers=[
|
||||
"Programming Language :: Python :: 3",
|
||||
"Environment :: GPU :: NVIDIA CUDA :: 12",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
],
|
||||
python_requires='>=3.10',
|
||||
install_requires=["torch>=2.5.0"])
|
||||
@@ -1,24 +0,0 @@
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/ATen.h>
|
||||
|
||||
#include <vector>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
|
||||
#ifdef TK_COMPILE_ATTN
|
||||
extern torch::Tensor sta_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text
|
||||
);
|
||||
#endif
|
||||
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "Sliding Block Attention Kernels"; // optional module docstring
|
||||
|
||||
#ifdef TK_COMPILE_ATTN
|
||||
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
|
||||
#endif
|
||||
}
|
||||
@@ -1,35 +0,0 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
from st_attn_cuda import sta_fwd
|
||||
|
||||
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True):
|
||||
seq_length = q_all.shape[2]
|
||||
if has_text:
|
||||
assert q_all.shape[
|
||||
2] >= 115200, "STA currently only supports video with latent size (30, 48, 80), which is 117 frames x 768 x 1280 pixels"
|
||||
assert q_all.shape[1] == len(window_size), "Number of heads must match the number of window sizes"
|
||||
target_size = math.ceil(seq_length / 384) * 384
|
||||
pad_size = target_size - seq_length
|
||||
if pad_size > 0:
|
||||
q_all = torch.cat([q_all, q_all[:, :, -pad_size:]], dim=2)
|
||||
k_all = torch.cat([k_all, k_all[:, :, -pad_size:]], dim=2)
|
||||
v_all = torch.cat([v_all, v_all[:, :, -pad_size:]], dim=2)
|
||||
else:
|
||||
assert q_all.shape[2] == 82944
|
||||
|
||||
hidden_states = torch.empty_like(q_all)
|
||||
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
|
||||
for head_index, (t_kernel, h_kernel, w_kernel) in enumerate(window_size):
|
||||
for batch in range(q_all.shape[0]):
|
||||
q_head, k_head, v_head, o_head = (q_all[batch:batch + 1, head_index:head_index + 1],
|
||||
k_all[batch:batch + 1,
|
||||
head_index:head_index + 1], v_all[batch:batch + 1,
|
||||
head_index:head_index + 1],
|
||||
hidden_states[batch:batch + 1, head_index:head_index + 1])
|
||||
|
||||
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text)
|
||||
if has_text:
|
||||
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True)
|
||||
return hidden_states[:, :, :seq_length]
|
||||
@@ -1,687 +0,0 @@
|
||||
// # Define TORCH_COMPILE macro
|
||||
|
||||
#include "kittens.cuh"
|
||||
#include <cooperative_groups.h>
|
||||
#include <iostream>
|
||||
#include <stdio.h>
|
||||
|
||||
#define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
|
||||
#define ABS(x) ((x) < 0 ? -(x) : (x))
|
||||
|
||||
constexpr int CONSUMER_WARPGROUPS = (3);
|
||||
constexpr int PRODUCER_WARPGROUPS = (1);
|
||||
constexpr int NUM_WARPGROUPS = (CONSUMER_WARPGROUPS+PRODUCER_WARPGROUPS);
|
||||
constexpr int NUM_WORKERS = (NUM_WARPGROUPS*kittens::WARPGROUP_WARPS);
|
||||
|
||||
using namespace kittens;
|
||||
namespace cg = cooperative_groups;
|
||||
|
||||
template<int D> struct fwd_attend_ker_tile_dims {};
|
||||
template<> struct fwd_attend_ker_tile_dims<64> {
|
||||
constexpr static int tile_width = (64);
|
||||
constexpr static int qo_height = (4*16);
|
||||
constexpr static int kv_height = (8*16);
|
||||
constexpr static int stages = (4);
|
||||
};
|
||||
template<> struct fwd_attend_ker_tile_dims<128> {
|
||||
constexpr static int tile_width = (128);
|
||||
constexpr static int qo_height = (4*16);
|
||||
constexpr static int kv_height = (8*16);
|
||||
constexpr static int stages = (2);
|
||||
};
|
||||
|
||||
template<int D> struct fwd_globals {
|
||||
using q_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
using k_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
using v_tile = st_bf<fwd_attend_ker_tile_dims<D>::kv_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>>;
|
||||
using o_tile = st_bf<fwd_attend_ker_tile_dims<D>::qo_height, fwd_attend_ker_tile_dims<D>::tile_width>;
|
||||
|
||||
using q_gl = gl<bf16, -1, -1, -1, -1, q_tile>;
|
||||
using k_gl = gl<bf16, -1, -1, -1, -1, k_tile>;
|
||||
using v_gl = gl<bf16, -1, -1, -1, -1, v_tile>;
|
||||
using l_gl = gl<float, -1, -1, -1, -1, l_col_vec>;
|
||||
using o_gl = gl<bf16, -1, -1, -1, -1, o_tile>;
|
||||
|
||||
q_gl q;
|
||||
k_gl k;
|
||||
v_gl v;
|
||||
l_gl l;
|
||||
o_gl o;
|
||||
|
||||
const int N;
|
||||
const int text_L;
|
||||
const int hr;
|
||||
};
|
||||
|
||||
|
||||
template<int D, bool is_causal, bool text_q, bool text_kv, int DT, int DH, int DW, int CT, int CH, int CW>
|
||||
__global__ __launch_bounds__((NUM_WORKERS)*kittens::WARP_THREADS, 1)
|
||||
void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
|
||||
extern __shared__ int __shm[];
|
||||
tma_swizzle_allocator al((int*)&__shm[0]);
|
||||
int warpid = kittens::warpid(), warpgroupid = warpid/kittens::WARPGROUP_WARPS;
|
||||
|
||||
using K = fwd_attend_ker_tile_dims<D>;
|
||||
|
||||
using q_tile = st_bf<K::qo_height, K::tile_width>;
|
||||
using k_tile = st_bf<K::kv_height, K::tile_width>;
|
||||
using v_tile = st_bf<K::kv_height, K::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<K::qo_height, K::tile_width>>;
|
||||
using o_tile = st_bf<K::qo_height, K::tile_width>;
|
||||
|
||||
q_tile (&q_smem)[CONSUMER_WARPGROUPS] = al.allocate<q_tile, CONSUMER_WARPGROUPS>();
|
||||
k_tile (&k_smem)[K::stages] = al.allocate<k_tile, K::stages >();
|
||||
v_tile (&v_smem)[K::stages] = al.allocate<v_tile, K::stages >();
|
||||
l_col_vec (&l_smem)[CONSUMER_WARPGROUPS] = al.allocate<l_col_vec, CONSUMER_WARPGROUPS>();
|
||||
auto (*o_smem) = reinterpret_cast<o_tile(*)>(q_smem);
|
||||
int img_kv_blocks;
|
||||
int kv_blocks = g.N / (K::kv_height);
|
||||
if constexpr (text_kv) {
|
||||
img_kv_blocks = kv_blocks - 3;
|
||||
} else {
|
||||
img_kv_blocks = kv_blocks;
|
||||
}
|
||||
int kv_head_idx = blockIdx.y / g.hr;
|
||||
int seq_idx;
|
||||
if constexpr (text_q) {
|
||||
seq_idx = CT * CH * CW * 6.0 + blockIdx.x * CONSUMER_WARPGROUPS;
|
||||
} else {
|
||||
seq_idx = blockIdx.x * CONSUMER_WARPGROUPS;
|
||||
}
|
||||
__shared__ kittens::semaphore qsmem_semaphore, k_smem_arrived[K::stages], v_smem_arrived[K::stages], compute_done[K::stages];
|
||||
if (threadIdx.x == 0) {
|
||||
init_semaphore(qsmem_semaphore, 0, 1);
|
||||
for(int j = 0; j < K::stages; j++) {
|
||||
init_semaphore(k_smem_arrived[j], 0, 1);
|
||||
init_semaphore(v_smem_arrived[j], 0, 1);
|
||||
init_semaphore(compute_done[j], CONSUMER_WARPGROUPS, 0);
|
||||
}
|
||||
|
||||
tma::expect_bytes(qsmem_semaphore, sizeof(q_smem));
|
||||
|
||||
for (int wg = 0; wg < CONSUMER_WARPGROUPS; wg++) {
|
||||
coord<q_tile> q_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + wg, 0};
|
||||
tma::load_async(q_smem[wg], g.q, q_tile_idx, qsmem_semaphore);
|
||||
}
|
||||
|
||||
if constexpr (text_q){
|
||||
for (int j = 0; j < K::stages - 1; j++) {
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
|
||||
tma::expect_bytes(k_smem_arrived[j], sizeof(k_tile));
|
||||
tma::load_async(k_smem[j], g.k, kv_tile_idx, k_smem_arrived[j]);
|
||||
tma::expect_bytes(v_smem_arrived[j], sizeof(v_tile));
|
||||
tma::load_async(v_smem[j], g.v, kv_tile_idx, v_smem_arrived[j]);
|
||||
}
|
||||
} else {
|
||||
int qt = seq_idx / 6 / (CH * CW);
|
||||
int qh = (seq_idx / 6) % (CH * CW) / CW;
|
||||
int qw = (seq_idx / 6) % CW;
|
||||
qt = CLAMP(qt, DT, CT-DT-1);
|
||||
qh = CLAMP(qh, DH, CH-DH-1);
|
||||
qw = CLAMP(qw, DW, CW-DW-1);
|
||||
int count = 0;
|
||||
int j = 0;
|
||||
while (count < K::stages - 1) {
|
||||
int kt = j / 3 / (CH * CW);
|
||||
int kh = (j / 3) % (CH * CW) / CW;
|
||||
int kw = (j / 3) % CW;
|
||||
bool mask = (ABS(qt - kt) <= DT) && (ABS(qh - kh) <= DH) && (ABS(qw - kw) <= DW);
|
||||
if (mask){
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
|
||||
tma::expect_bytes(k_smem_arrived[count], sizeof(k_tile));
|
||||
tma::load_async(k_smem[count], g.k, kv_tile_idx, k_smem_arrived[count]);
|
||||
tma::expect_bytes(v_smem_arrived[count], sizeof(v_tile));
|
||||
tma::load_async(v_smem[count], g.v, kv_tile_idx, v_smem_arrived[count]);
|
||||
count += 1;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
int pipe_idx = K::stages - 1;
|
||||
|
||||
if(warpgroupid == NUM_WARPGROUPS-1) {
|
||||
warpgroup::decrease_registers<32>();
|
||||
|
||||
int kv_iters;
|
||||
if constexpr (is_causal) {
|
||||
kv_iters = (seq_idx * (K::qo_height/kittens::TILE_ROW_DIM<bf16>)) - 1 + (CONSUMER_WARPGROUPS * (K::qo_height/kittens::TILE_ROW_DIM<bf16>));
|
||||
kv_iters = ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) == 0) ? (0) : ((kv_iters / (K::kv_height/kittens::TILE_ROW_DIM<bf16>)) - 1);
|
||||
}
|
||||
else { kv_iters = kv_blocks-2;}
|
||||
|
||||
if(warpid == NUM_WORKERS-4) {
|
||||
if constexpr (text_q){
|
||||
for (auto kv_idx = pipe_idx - 1; kv_idx <= kv_iters; kv_idx++) {
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, kv_idx + 1, 0};
|
||||
tma::expect_bytes(k_smem_arrived[(kv_idx+1)%K::stages], sizeof(k_tile));
|
||||
tma::load_async(k_smem[(kv_idx+1)%K::stages], g.k, kv_tile_idx, k_smem_arrived[(kv_idx+1)%K::stages]);
|
||||
tma::expect_bytes(v_smem_arrived[(kv_idx+1)%K::stages], sizeof(v_tile));
|
||||
tma::load_async(v_smem[(kv_idx+1)%K::stages], g.v, kv_tile_idx, v_smem_arrived[(kv_idx+1)%K::stages]);
|
||||
kittens::wait(compute_done[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
}
|
||||
} else {
|
||||
int qt = seq_idx / 6 / (CH * CW);
|
||||
int qh = (seq_idx / 6) % (CH * CW) / CW;
|
||||
int qw = (seq_idx / 6) % CW;
|
||||
qt = CLAMP(qt, DT, CT-DT-1);
|
||||
qh = CLAMP(qh, DH, CH-DH-1);
|
||||
qw = CLAMP(qw, DW, CW-DW-1);
|
||||
int k_t_min = CLAMP(qt-DT, 0, CT-1);
|
||||
int k_t_max = CLAMP(qt+DT, 0, CT-1);
|
||||
int k_h_min = CLAMP(qh-DH, 0, CH-1);
|
||||
int k_h_max = CLAMP(qh+DH, 0, CH-1);
|
||||
int k_w_min = CLAMP(qw-DW, 0, CW-1);
|
||||
int k_w_max = CLAMP(qw+DW, 0, CW-1);
|
||||
int count = 0;
|
||||
for (int kt = k_t_min; kt <= k_t_max; kt++) {
|
||||
for (int kh = k_h_min; kh <= k_h_max; kh++) {
|
||||
for (int kw = k_w_min; kw <= k_w_max; kw++) {
|
||||
for (int j = 0; j <= 2; j++){
|
||||
if (count >= K::stages - 1) {
|
||||
int index = ((kt * (CH * CW)) + (kh * CW) + kw) * 3 + j;
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
|
||||
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
|
||||
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
|
||||
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
|
||||
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
|
||||
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
|
||||
count += 1;
|
||||
} else {
|
||||
count += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
// for text
|
||||
for (int index = img_kv_blocks; index < kv_blocks; index++) {
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, index, 0};
|
||||
tma::expect_bytes(k_smem_arrived[count%K::stages], sizeof(k_tile));
|
||||
tma::load_async(k_smem[count%K::stages], g.k, kv_tile_idx, k_smem_arrived[count%K::stages]);
|
||||
tma::expect_bytes(v_smem_arrived[count%K::stages], sizeof(v_tile));
|
||||
tma::load_async(v_smem[count%K::stages], g.v, kv_tile_idx, v_smem_arrived[count%K::stages]);
|
||||
kittens::wait(compute_done[(count - 1)%K::stages], ((count - 1)/K::stages)%2);
|
||||
count += 1;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
}
|
||||
}
|
||||
else {
|
||||
warpgroup::increase_registers<160>();
|
||||
|
||||
rt_fl<16, K::kv_height> att_block;
|
||||
rt_bf<16, K::kv_height> att_block_mma;
|
||||
rt_fl<16, K::tile_width> o_reg;
|
||||
|
||||
col_vec<rt_fl<16, K::kv_height>> max_vec, norm_vec, max_vec_last_scaled, max_vec_scaled;
|
||||
|
||||
neg_infty(max_vec);
|
||||
zero(norm_vec);
|
||||
zero(o_reg);
|
||||
|
||||
int kv_iters;
|
||||
if constexpr (is_causal) {
|
||||
kv_iters = (seq_idx * 4) - 1 + (CONSUMER_WARPGROUPS * 4);
|
||||
kv_iters = (kv_iters/8);
|
||||
}
|
||||
else if constexpr (text_q){
|
||||
// the last three kv blocks are for text, we process them separately
|
||||
kv_iters = img_kv_blocks - 1;
|
||||
} else {
|
||||
kv_iters = CLAMP(DT*2+1, 1, CT) * CLAMP(DH*2+1, 1, CH) * CLAMP(DW*2+1, 1, CW) * 3 - 1 ;
|
||||
}
|
||||
|
||||
kittens::wait(qsmem_semaphore, 0);
|
||||
for (auto kv_idx = 0; kv_idx <= kv_iters; kv_idx++) {
|
||||
|
||||
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
|
||||
|
||||
copy(max_vec_last_scaled, max_vec);
|
||||
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
|
||||
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
|
||||
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
row_max(max_vec, att_block, max_vec);
|
||||
|
||||
if constexpr (D == 64) {
|
||||
mul(att_block, att_block, 1.44269504089f*0.125f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
|
||||
}
|
||||
else {
|
||||
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
|
||||
}
|
||||
|
||||
sub_row(att_block, att_block, max_vec_scaled);
|
||||
exp2(att_block, att_block);
|
||||
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
|
||||
exp2(max_vec_last_scaled, max_vec_last_scaled);
|
||||
mul(norm_vec, norm_vec, max_vec_last_scaled);
|
||||
row_sum(norm_vec, att_block, norm_vec);
|
||||
add(att_block, att_block, 0.f);
|
||||
copy(att_block_mma, att_block);
|
||||
mul_row(o_reg, o_reg, max_vec_last_scaled);
|
||||
|
||||
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
|
||||
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
|
||||
}
|
||||
// the last three kv blocks are for text, we process them separately
|
||||
if constexpr(text_kv) {
|
||||
for (auto kv_idx = kv_iters + 1; kv_idx <= kv_iters + 3; kv_idx++) {
|
||||
|
||||
kittens::wait(k_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
warpgroup::mm_ABt(att_block, q_smem[warpgroupid], k_smem[(kv_idx)%K::stages]);
|
||||
|
||||
copy(max_vec_last_scaled, max_vec);
|
||||
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
|
||||
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
|
||||
|
||||
warpgroup::mma_async_wait();
|
||||
// apply non-pad mask
|
||||
int offset = g.text_L - (kv_idx - (kv_iters + 1)) * K::kv_height;
|
||||
// printf("k_idx_start: %d, k_idx_end: %d, text_end: %d, offset: %d\n", k_idx_start, k_idx_end, text_end, offset);
|
||||
right_fill(att_block, att_block, offset, base_types::constants<float>::neg_infty());
|
||||
|
||||
|
||||
row_max(max_vec, att_block, max_vec);
|
||||
|
||||
if constexpr (D == 64) {
|
||||
mul(att_block, att_block, 1.44269504089f*0.125f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
|
||||
}
|
||||
else {
|
||||
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
|
||||
}
|
||||
|
||||
sub_row(att_block, att_block, max_vec_scaled);
|
||||
exp2(att_block, att_block);
|
||||
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
|
||||
exp2(max_vec_last_scaled, max_vec_last_scaled);
|
||||
mul(norm_vec, norm_vec, max_vec_last_scaled);
|
||||
row_sum(norm_vec, att_block, norm_vec);
|
||||
add(att_block, att_block, 0.f);
|
||||
copy(att_block_mma, att_block);
|
||||
mul_row(o_reg, o_reg, max_vec_last_scaled);
|
||||
|
||||
kittens::wait(v_smem_arrived[(kv_idx)%K::stages], (kv_idx/K::stages)%2);
|
||||
|
||||
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[(kv_idx)%K::stages]);
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
if(warpgroup::laneid() == 0) arrive(compute_done[(kv_idx)%K::stages], 1);
|
||||
}
|
||||
}
|
||||
|
||||
div_row(o_reg, o_reg, norm_vec);
|
||||
warpgroup::store(o_smem[warpgroupid], o_reg);
|
||||
warpgroup::sync(warpgroupid+4);
|
||||
|
||||
if (warpid % 4 == 0) {
|
||||
coord<o_tile> o_tile_idx = {blockIdx.z, blockIdx.y, (seq_idx) + warpgroupid, 0};
|
||||
tma::store_async(g.o, o_smem[warpgroupid], o_tile_idx);
|
||||
}
|
||||
|
||||
mul(max_vec_scaled, max_vec_scaled, 0.69314718056f);
|
||||
log(norm_vec, norm_vec);
|
||||
add(norm_vec, norm_vec, max_vec_scaled);
|
||||
|
||||
if constexpr (D == 64) { mul(norm_vec, norm_vec, -8.0f); }
|
||||
else { mul(norm_vec, norm_vec, -11.313708499f); }
|
||||
|
||||
warpgroup::store(l_smem[warpgroupid], norm_vec);
|
||||
warpgroup::sync(warpgroupid+4);
|
||||
|
||||
if (warpid % 4 == 0) {
|
||||
coord<l_col_vec> tile_idx = {blockIdx.z, blockIdx.y, 0, (seq_idx) + warpgroupid};
|
||||
tma::store_async(g.l, l_smem[warpgroupid], tile_idx);
|
||||
}
|
||||
tma::store_async_wait();
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
#include "pyutils/torch_helpers.cuh"
|
||||
#include <ATen/cuda/CUDAContext.h>
|
||||
#include <iostream>
|
||||
|
||||
torch::Tensor
|
||||
sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_h_size, int kernel_w_size, int text_length, bool process_text, bool has_text)
|
||||
{
|
||||
CHECK_INPUT(q);
|
||||
CHECK_INPUT(k);
|
||||
CHECK_INPUT(v);
|
||||
|
||||
auto batch = q.size(0);
|
||||
auto seq_len = q.size(2);
|
||||
auto head_dim = q.size(3);
|
||||
auto qo_heads = q.size(1);
|
||||
auto kv_heads = k.size(1);
|
||||
|
||||
// check to see that these dimensions match for all inputs
|
||||
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(0) == batch, "V batch dimension - idx 0 - must match for all inputs");
|
||||
|
||||
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
|
||||
|
||||
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
|
||||
TORCH_CHECK(v.size(3) == head_dim, "V head dimension - idx 3 - must match for all non-vector inputs");
|
||||
|
||||
TORCH_CHECK(qo_heads >= kv_heads, "QO heads must be greater than or equal to KV heads");
|
||||
TORCH_CHECK(qo_heads % kv_heads == 0, "QO heads must be divisible by KV heads");
|
||||
TORCH_CHECK(q.size(1) == qo_heads, "QO head dimension - idx 1 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
|
||||
TORCH_CHECK(v.size(1) == kv_heads, "KV head dimension - idx 1 - must match for all inputs");
|
||||
|
||||
auto hr = qo_heads / kv_heads;
|
||||
|
||||
c10::BFloat16* q_ptr = q.data_ptr<c10::BFloat16>();
|
||||
c10::BFloat16* k_ptr = k.data_ptr<c10::BFloat16>();
|
||||
c10::BFloat16* v_ptr = v.data_ptr<c10::BFloat16>();
|
||||
|
||||
bf16* d_q = reinterpret_cast<bf16*>(q_ptr);
|
||||
bf16* d_k = reinterpret_cast<bf16*>(k_ptr);
|
||||
bf16* d_v = reinterpret_cast<bf16*>(v_ptr);
|
||||
|
||||
|
||||
|
||||
torch::Tensor l_vec = torch::empty({static_cast<const uint>(batch),
|
||||
static_cast<const uint>(qo_heads),
|
||||
static_cast<const uint>(seq_len),
|
||||
static_cast<const uint>(1)},
|
||||
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
|
||||
|
||||
|
||||
bf16* o_ptr = reinterpret_cast<bf16*>(o.data_ptr<c10::BFloat16>());
|
||||
bf16* d_o = reinterpret_cast<bf16*>(o_ptr);
|
||||
|
||||
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
|
||||
float* d_l = reinterpret_cast<float*>(l_ptr);
|
||||
|
||||
cudaDeviceSynchronize();
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
|
||||
if (head_dim == 128) {
|
||||
using q_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using k_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using v_tile = st_bf<fwd_attend_ker_tile_dims<128>::kv_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>>;
|
||||
using o_tile = st_bf<fwd_attend_ker_tile_dims<128>::qo_height, fwd_attend_ker_tile_dims<128>::tile_width>;
|
||||
|
||||
using q_global = gl<bf16, -1, -1, -1, -1, q_tile>;
|
||||
using k_global = gl<bf16, -1, -1, -1, -1, k_tile>;
|
||||
using v_global = gl<bf16, -1, -1, -1, -1, v_tile>;
|
||||
using l_global = gl<float, -1, -1, -1, -1, l_col_vec>;
|
||||
using o_global = gl<bf16, -1, -1, -1, -1, o_tile>;
|
||||
|
||||
using globals = fwd_globals<128>;
|
||||
|
||||
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
|
||||
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(text_length), static_cast<int>(hr)};
|
||||
|
||||
auto mem_size = kittens::MAX_SHARED_MEMORY;
|
||||
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_text(2, qo_heads, batch);
|
||||
if (!process_text) {
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 1, 1, 1, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 1, 1, 1, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 1, 1, 2, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true,1, 1, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 5 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 1, 1, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 2, 1, 1, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
}else if (kernel_t_size ==3 && kernel_h_size == 5 && kernel_w_size == 5){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 1, 2, 2, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 1, 2, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size ==5 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 3, 0, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 2, 3, 0, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size ==5 && kernel_h_size == 3 && kernel_w_size == 5){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 1, 2, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 2, 1, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 5){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 2, 2, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 2, 2, 2, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 5 && kernel_h_size == 5 && kernel_w_size == 7){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 2, 3, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 2, 2, 3, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 5 && kernel_h_size == 6 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 3, 5, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 2, 3, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 0, 0, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 2, 0, 0, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 0, 3, 5, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true, 0, 3, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 5 && kernel_h_size == 1 && kernel_w_size == 10){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, true, 2, 0, 5, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, true,2, 0, 5, 5, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else {
|
||||
// print error
|
||||
std::cout << "Invalid kernel size" << std::endl;
|
||||
//print kernel size
|
||||
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
|
||||
}
|
||||
} else {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, true, true, 1, 1, 1, 5, 6, 10><<<grid_text, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
}
|
||||
|
||||
} else {
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
|
||||
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 6) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 1, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false,1, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 3 && kernel_w_size == 3) {
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
}else if (kernel_t_size ==3 && kernel_h_size == 6 && kernel_w_size == 3){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 1, 3, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size ==6 && kernel_h_size == 3 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 1, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 1 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 0, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 6 && kernel_w_size == 1){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 3, 0, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else if (kernel_t_size == 6 && kernel_h_size == 1 && kernel_w_size == 6){
|
||||
cudaFuncSetAttribute(
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
mem_size
|
||||
);
|
||||
fwd_attend_ker<128, false, false, false, 3, 0, 3, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
|
||||
} else {
|
||||
// print error
|
||||
std::cout << "Invalid kernel size" << std::endl;
|
||||
//print kernel size
|
||||
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
|
||||
}
|
||||
|
||||
}
|
||||
CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
cudaStreamSynchronize(stream);
|
||||
}
|
||||
|
||||
return o;
|
||||
cudaDeviceSynchronize();
|
||||
}
|
||||
@@ -1,151 +0,0 @@
|
||||
import os
|
||||
from collections import defaultdict
|
||||
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import torch
|
||||
from st_attn import sliding_tile_attention
|
||||
|
||||
|
||||
def flops(batch, seqlen, nheads, headdim, causal, mode="fwd"):
|
||||
assert mode in ["fwd", "bwd", "fwd_bwd"]
|
||||
f = 4 * batch * seqlen**2 * nheads * headdim // (2 if causal else 1)
|
||||
return f if mode == "fwd" else (2.5 * f if mode == "bwd" else 3.5 * f)
|
||||
|
||||
|
||||
def efficiency(flop, time):
|
||||
flop = flop / 1e12
|
||||
time = time / 1e6
|
||||
return flop / time
|
||||
|
||||
|
||||
def benchmark_attention(configurations):
|
||||
results = {'fwd': defaultdict(list), 'bwd': defaultdict(list)}
|
||||
|
||||
for B, H, N, D, causal in configurations:
|
||||
print("=" * 60)
|
||||
print(f"Timing forward and backward pass for B={B}, H={H}, N={N}, D={D}, causal={causal}")
|
||||
|
||||
q = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
|
||||
k = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
|
||||
v = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
|
||||
|
||||
grad_output = torch.randn_like(q, requires_grad=False).contiguous()
|
||||
|
||||
qg = torch.zeros_like(q, requires_grad=False, dtype=torch.float).contiguous()
|
||||
kg = torch.zeros_like(k, requires_grad=False, dtype=torch.float).contiguous()
|
||||
vg = torch.zeros_like(v, requires_grad=False, dtype=torch.float).contiguous()
|
||||
|
||||
# Prepare for timing forward pass
|
||||
start_events_fwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
|
||||
end_events_fwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Warmup for forward pass
|
||||
for _ in range(10):
|
||||
o = sliding_tile_attention(q, k, v, [[6, 6, 6]] * 24, 0, False)
|
||||
|
||||
# Time the forward pass
|
||||
for i in range(10):
|
||||
start_events_fwd[i].record()
|
||||
o = sliding_tile_attention(q, k, v, [[6, 6, 6]] * 24, 0, False)
|
||||
end_events_fwd[i].record()
|
||||
|
||||
torch.cuda.synchronize()
|
||||
times_fwd = [s.elapsed_time(e) for s, e in zip(start_events_fwd, end_events_fwd)]
|
||||
time_us_fwd = np.mean(times_fwd) * 1000
|
||||
|
||||
tflops_fwd = efficiency(flops(B, N, H, D, causal, 'fwd'), time_us_fwd)
|
||||
results['fwd'][(D, causal)].append((N, tflops_fwd))
|
||||
|
||||
print(f"Average time for forward pass in us: {time_us_fwd:.2f}")
|
||||
print(f"Average efficiency for forward pass in TFLOPS: {tflops_fwd}")
|
||||
print("-" * 60)
|
||||
|
||||
# torch.cuda.empty_cache()
|
||||
# torch.cuda.synchronize()
|
||||
|
||||
# # Prepare for timing backward pass
|
||||
# start_events_bwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
|
||||
# end_events_bwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
|
||||
|
||||
# # Warmup for backward pass
|
||||
# for _ in range(10):
|
||||
# qg, kg, vg = tk.mha_backward(q, k, v, o, l_vec, grad_output, causal)
|
||||
|
||||
# # Time the backward pass
|
||||
# for i in range(10):
|
||||
# start_events_bwd[i].record()
|
||||
# qg, kg, vg = tk.mha_backward(q, k, v, o, l_vec, grad_output, causal)
|
||||
# end_events_bwd[i].record()
|
||||
|
||||
# torch.cuda.synchronize()
|
||||
# times_bwd = [s.elapsed_time(e) for s, e in zip(start_events_bwd, end_events_bwd)]
|
||||
# time_us_bwd = np.mean(times_bwd) * 1000
|
||||
|
||||
# tflops_bwd = efficiency(flops(B, N, H, D, causal, 'bwd'), time_us_bwd)
|
||||
# results['bwd'][(D, causal)].append((N, tflops_bwd))
|
||||
|
||||
# print(f"Average time for backward pass in us: {time_us_bwd:.2f}")
|
||||
# print(f"Average efficiency for backward pass in TFLOPS: {tflops_bwd}")
|
||||
print("=" * 60)
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def plot_results(results):
|
||||
os.makedirs('benchmark_results', exist_ok=True)
|
||||
for mode in ['fwd', 'bwd']:
|
||||
for (D, causal), values in results[mode].items():
|
||||
seq_lens = [x[0] for x in values]
|
||||
tflops = [x[1] for x in values]
|
||||
|
||||
plt.figure(figsize=(10, 6))
|
||||
bars = plt.bar(range(len(seq_lens)), tflops, tick_label=seq_lens)
|
||||
plt.xlabel('Sequence Length')
|
||||
plt.ylabel('TFLOPS')
|
||||
plt.title(f'{mode.upper()} Pass - Head Dim: {D}, Causal: {causal}')
|
||||
plt.grid(True)
|
||||
|
||||
# Adding the numerical y value on top of each bar
|
||||
for bar in bars:
|
||||
yval = bar.get_height()
|
||||
plt.text(bar.get_x() + bar.get_width() / 2, yval, round(yval, 2), ha='center', va='bottom')
|
||||
|
||||
filename = f'benchmark_results/{mode}_D{D}_causal{causal}.png'
|
||||
plt.savefig(filename)
|
||||
plt.close()
|
||||
|
||||
|
||||
# Example list of configurations to test
|
||||
configurations = [
|
||||
(2, 24, 82944, 128, False),
|
||||
# (16, 16, 768*16, 128, False),
|
||||
# (16, 16, 768*2, 128, False),
|
||||
# (16, 16, 768*4, 128, False),
|
||||
# (16, 16, 768*8, 128, False),
|
||||
# (16, 16, 768*16, 128, False),
|
||||
# (16, 16, 768, 128, True),
|
||||
# (16, 16, 768*2, 128, True),
|
||||
# (16, 16, 768*4, 128, True),
|
||||
# (16, 16, 768*8, 128, True),
|
||||
# (16, 16, 768*16, 128, True),
|
||||
# (16, 32, 768, 64, False),
|
||||
# (16, 32, 768*2, 64, False),
|
||||
# (16, 32, 768*4, 64, False),
|
||||
# (16, 32, 768*8, 64, False),
|
||||
# (16, 32, 768*16, 64, False),
|
||||
# (16, 32, 768, 64, True),
|
||||
# (16, 32, 768*2, 64, True),
|
||||
# (16, 32, 768*4, 64, True),
|
||||
# (16, 32, 768*8, 64, True),
|
||||
# (16, 32, 768*16, 64, True),
|
||||
]
|
||||
|
||||
results = benchmark_attention(configurations)
|
||||
# plot_results(results)
|
||||
@@ -1,71 +0,0 @@
|
||||
from typing import Tuple
|
||||
|
||||
import torch
|
||||
from torch import BoolTensor, IntTensor
|
||||
from torch.nn.attention.flex_attention import create_block_mask
|
||||
|
||||
# Peiyuan: This is neccesay. Dont know why. see https://github.com/pytorch/pytorch/issues/135028
|
||||
torch._inductor.config.realize_opcount_threshold = 100
|
||||
|
||||
|
||||
def generate_sta_mask(canvas_twh, kernel_twh, tile_twh, text_length):
|
||||
"""Generates a 3D NATTEN attention mask with a given kernel size.
|
||||
|
||||
Args:
|
||||
canvas_t: The time dimension of the canvas.
|
||||
canvas_h: The height of the canvas.
|
||||
canvas_w: The width of the canvas.
|
||||
kernel_t: The time dimension of the kernel.
|
||||
kernel_h: The height of the kernel.
|
||||
kernel_w: The width of the kernel.
|
||||
"""
|
||||
canvas_t, canvas_h, canvas_w = canvas_twh
|
||||
kernel_t, kernel_h, kernel_w = kernel_twh
|
||||
tile_t_size, tile_h_size, tile_w_size = tile_twh
|
||||
total_tile_size = tile_t_size * tile_h_size * tile_w_size
|
||||
canvas_tile_t, canvas_tile_h, canvas_tile_w = canvas_t // tile_t_size, canvas_h // tile_h_size, canvas_w // tile_w_size
|
||||
img_seq_len = canvas_t * canvas_h * canvas_w
|
||||
|
||||
def get_tile_t_x_y(idx: IntTensor) -> Tuple[IntTensor, IntTensor, IntTensor]:
|
||||
tile_id = idx // total_tile_size
|
||||
tile_t = tile_id // (canvas_tile_h * canvas_tile_w)
|
||||
tile_h = (tile_id % (canvas_tile_h * canvas_tile_w)) // canvas_tile_w
|
||||
tile_w = tile_id % canvas_tile_w
|
||||
return tile_t, tile_h, tile_w
|
||||
|
||||
def sta_mask_3d(
|
||||
b: IntTensor,
|
||||
h: IntTensor,
|
||||
q_idx: IntTensor,
|
||||
kv_idx: IntTensor,
|
||||
) -> BoolTensor:
|
||||
q_t_tile, q_x_tile, q_y_tile = get_tile_t_x_y(q_idx)
|
||||
kv_t_tile, kv_x_tile, kv_y_tile = get_tile_t_x_y(kv_idx)
|
||||
# kernel nominally attempts to center itself on the query, but kernel center
|
||||
# is clamped to a fixed distance (kernel half-length) from the canvas edge
|
||||
kernel_center_t = q_t_tile.clamp(kernel_t // 2, (canvas_tile_t - 1) - kernel_t // 2)
|
||||
kernel_center_x = q_x_tile.clamp(kernel_h // 2, (canvas_tile_h - 1) - kernel_h // 2)
|
||||
kernel_center_y = q_y_tile.clamp(kernel_w // 2, (canvas_tile_w - 1) - kernel_w // 2)
|
||||
time_mask = (kernel_center_t - kv_t_tile).abs() <= kernel_t // 2
|
||||
hori_mask = (kernel_center_x - kv_x_tile).abs() <= kernel_h // 2
|
||||
vert_mask = (kernel_center_y - kv_y_tile).abs() <= kernel_w // 2
|
||||
image_mask = (q_idx < img_seq_len) & (kv_idx < img_seq_len)
|
||||
image_to_text_mask = (q_idx < img_seq_len) & (kv_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
|
||||
text_to_all_mask = (q_idx >= img_seq_len) & (kv_idx < img_seq_len + text_length)
|
||||
return (image_mask & time_mask & hori_mask & vert_mask) | image_to_text_mask | text_to_all_mask
|
||||
|
||||
sta_mask_3d.__name__ = f"natten_3d_c{canvas_t}x{canvas_w}x{canvas_h}_k{kernel_t}x{kernel_w}x{kernel_h}"
|
||||
return sta_mask_3d
|
||||
|
||||
|
||||
def get_sliding_tile_attention_mask(kernel_size, tile_size, img_size, text_length, device, text_max_len=256):
|
||||
img_seq_len = img_size[0] * img_size[1] * img_size[2]
|
||||
image_mask = generate_sta_mask(img_size, kernel_size, tile_size, text_length)
|
||||
mask = create_block_mask(image_mask,
|
||||
B=None,
|
||||
H=None,
|
||||
Q_LEN=img_seq_len + text_max_len,
|
||||
KV_LEN=img_seq_len + text_max_len,
|
||||
device=device,
|
||||
_compile=True)
|
||||
return mask
|
||||
@@ -1,96 +0,0 @@
|
||||
import torch
|
||||
from flex_sta_ref import get_sliding_tile_attention_mask
|
||||
from st_attn import sliding_tile_attention
|
||||
from torch.nn.attention.flex_attention import flex_attention
|
||||
# from flash_attn_interface import flash_attn_func
|
||||
from tqdm import tqdm
|
||||
|
||||
flex_attention = torch.compile(flex_attention, dynamic=False)
|
||||
|
||||
|
||||
def flex_test(Q, K, V, kernel_size):
|
||||
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (36, 48, 48), 39, 'cuda', 0)
|
||||
output = flex_attention(Q, K, V, block_mask=mask)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def h100_fwd_kernel_test(Q, K, V, kernel_size):
|
||||
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 39, False)
|
||||
return o
|
||||
|
||||
|
||||
def generate_tensor(shape, mean, std, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
|
||||
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
|
||||
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
|
||||
|
||||
return scaled_tensor.contiguous()
|
||||
|
||||
|
||||
def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mode='all'):
|
||||
results = {
|
||||
'TK vs FLEX': {
|
||||
'sum_diff': 0,
|
||||
'sum_abs': 0,
|
||||
'max_diff': 0
|
||||
},
|
||||
}
|
||||
kernel_size_ls = [(6, 1, 6), (6, 6, 1)]
|
||||
from tqdm import tqdm
|
||||
for kernel_size in tqdm(kernel_size_ls):
|
||||
for _ in range(num_iterations):
|
||||
torch.manual_seed(0)
|
||||
|
||||
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
tk_o = h100_fwd_kernel_test(Q, K, V, kernel_size)
|
||||
pt_o = flex_test(Q, K, V, kernel_size)
|
||||
|
||||
diff = pt_o - tk_o
|
||||
abs_diff = torch.abs(diff)
|
||||
results['TK vs FLEX']['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results['TK vs FLEX']['max_diff'] = max(results['TK vs FLEX']['max_diff'], torch.max(abs_diff).item())
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
print("kernel_size", kernel_size)
|
||||
print("max_diff", torch.max(abs_diff).item())
|
||||
print(
|
||||
"avg_diff",
|
||||
torch.sum(abs_diff).item() / (b * h * n * d *
|
||||
(1 if error_mode == 'output' else 3 if error_mode == 'backward' else 4)))
|
||||
|
||||
total_elements = b * h * n * d * num_iterations * (1 if error_mode == 'output' else
|
||||
3 if error_mode == 'backward' else 4) * len(kernel_size_ls)
|
||||
for name, data in results.items():
|
||||
avg_diff = data['sum_diff'] / total_elements
|
||||
max_diff = data['max_diff']
|
||||
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
|
||||
|
||||
return results
|
||||
|
||||
|
||||
def generate_error_graphs(b, h, d, causal, mean, std, error_mode='all'):
|
||||
seq_lengths = [82944]
|
||||
|
||||
tk_avg_errors, tk_max_errors = [], []
|
||||
|
||||
for n in tqdm(seq_lengths, desc="Generating error data"):
|
||||
results = check_correctness(b, h, n, d, causal, mean, std, error_mode=error_mode)
|
||||
|
||||
tk_avg_errors.append(results['TK vs FLEX']['avg_diff'])
|
||||
tk_max_errors.append(results['TK vs FLEX']['max_diff'])
|
||||
|
||||
|
||||
# Example usage
|
||||
b, h, d = 2, 24, 128
|
||||
causal = False
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
for mode in ['output']:
|
||||
generate_error_graphs(b, h, d, causal, mean, std, error_mode=mode)
|
||||
|
||||
print("Error graphs generated and saved for all modes.")
|
||||
Submodule csrc/sliding_tile_attention/tk deleted from 1719fb7264
+23
-14
@@ -1,15 +1,13 @@
|
||||
import argparse
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import gradio as gr
|
||||
import torch
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import export_to_video
|
||||
|
||||
from fastvideo.distill.solver import PCMFMScheduler
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
import tempfile
|
||||
import os
|
||||
import argparse
|
||||
|
||||
|
||||
def init_args():
|
||||
@@ -34,6 +32,7 @@ def init_args():
|
||||
|
||||
|
||||
def load_model(args):
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
if args.scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
else:
|
||||
@@ -50,9 +49,13 @@ def load_model(args):
|
||||
if args.transformer_path:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
|
||||
else:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder="transformer/")
|
||||
transformer = MochiTransformer3DModel.from_pretrained(
|
||||
args.model_path, subfolder="transformer/"
|
||||
)
|
||||
|
||||
pipe = MochiPipeline.from_pretrained(args.model_path, transformer=transformer, scheduler=scheduler)
|
||||
pipe = MochiPipeline.from_pretrained(
|
||||
args.model_path, transformer=transformer, scheduler=scheduler
|
||||
)
|
||||
pipe.enable_vae_tiling()
|
||||
# pipe.to(device)
|
||||
# if args.cpu_offload:
|
||||
@@ -73,7 +76,7 @@ def generate_video(
|
||||
randomize_seed=False,
|
||||
):
|
||||
if randomize_seed:
|
||||
seed = torch.randint(0, 1000000, (1, )).item()
|
||||
seed = torch.randint(0, 1000000, (1,)).item()
|
||||
|
||||
generator = torch.Generator(device="cuda").manual_seed(seed)
|
||||
|
||||
@@ -131,7 +134,9 @@ with gr.Blocks() as demo:
|
||||
step=32,
|
||||
value=args.height,
|
||||
)
|
||||
width = gr.Slider(label="Width", minimum=256, maximum=1024, step=32, value=args.width)
|
||||
width = gr.Slider(
|
||||
label="Width", minimum=256, maximum=1024, step=32, value=args.width
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
num_frames = gr.Slider(
|
||||
@@ -154,7 +159,9 @@ with gr.Blocks() as demo:
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
use_negative_prompt = gr.Checkbox(label="Use negative prompt", value=False)
|
||||
use_negative_prompt = gr.Checkbox(
|
||||
label="Use negative prompt", value=False
|
||||
)
|
||||
negative_prompt = gr.Text(
|
||||
label="Negative prompt",
|
||||
max_lines=1,
|
||||
@@ -162,7 +169,9 @@ with gr.Blocks() as demo:
|
||||
visible=False,
|
||||
)
|
||||
|
||||
seed = gr.Slider(label="Seed", minimum=0, maximum=1000000, step=1, value=args.seed)
|
||||
seed = gr.Slider(
|
||||
label="Seed", minimum=0, maximum=1000000, step=1, value=args.seed
|
||||
)
|
||||
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
|
||||
seed_output = gr.Number(label="Used Seed")
|
||||
|
||||
@@ -192,4 +201,4 @@ with gr.Blocks() as demo:
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
|
||||
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
|
||||
@@ -1,15 +0,0 @@
|
||||
Fast-Hunyuan comparison with original Hunyuan, achieving an 8X diffusion speed boost with the FastVideo framework.
|
||||
|
||||
https://github.com/user-attachments/assets/064ac1d2-11ed-4a0c-955b-4d412a96ef30
|
||||
|
||||
Fast-Mochi comparison with original Mochi, achieving an 8X diffusion speed boost with the FastVideo framework.
|
||||
|
||||
https://github.com/user-attachments/assets/5fbc4596-56d6-43aa-98e0-da472cf8e26c
|
||||
|
||||
Comparison between OpenAI Sora, original Hunyuan and FastHunyuan
|
||||
|
||||
https://github.com/user-attachments/assets/d323b712-3f68-42b2-952b-94f6a49c4836
|
||||
|
||||
Comparison between original FastHunyuan, LLM-INT8 quantized FastHunyuan and NF4 quantized FastHunyuan
|
||||
|
||||
https://github.com/user-attachments/assets/cf89efb5-5f68-4949-a085-f41c1ef26c94
|
||||
@@ -1,24 +0,0 @@
|
||||
# Minimal makefile for Sphinx documentation
|
||||
#
|
||||
|
||||
# You can set these variables from the command line, and also
|
||||
# from the environment for the first two.
|
||||
SPHINXOPTS ?=
|
||||
SPHINXBUILD ?= sphinx-build
|
||||
SOURCEDIR = source
|
||||
BUILDDIR = build
|
||||
|
||||
# Put it first so that "make" without argument is like "make help".
|
||||
help:
|
||||
@$(SPHINXBUILD) -M help "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
|
||||
|
||||
.PHONY: help Makefile
|
||||
|
||||
# Catch-all target: route all unknown targets to Sphinx using the new
|
||||
# "make mode" option. $(O) is meant as a shortcut for $(SPHINXOPTS).
|
||||
%: Makefile
|
||||
@$(SPHINXBUILD) -M $@ "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
|
||||
|
||||
clean:
|
||||
@$(SPHINXBUILD) -M clean "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
|
||||
rm -rf "$(SOURCEDIR)/getting_started/examples"
|
||||
@@ -1,20 +0,0 @@
|
||||
# FastVideo documents
|
||||
|
||||
## Build the docs
|
||||
|
||||
```bash
|
||||
# Install dependencies.
|
||||
pip install -r requirements-docs.txt
|
||||
|
||||
# Build the docs.
|
||||
make clean
|
||||
make html
|
||||
```
|
||||
|
||||
## Open the docs with your browser
|
||||
|
||||
```bash
|
||||
python -m http.server -d build/html/
|
||||
```
|
||||
|
||||
Launch your browser and open localhost:8000.
|
||||
@@ -1,16 +1,16 @@
|
||||
|
||||
|
||||
|
||||
## 🧱 Data Preprocess
|
||||
|
||||
To save GPU memory, we precompute text embeddings and VAE latents to eliminate the need to load the text encoder and VAE during training.
|
||||
|
||||
We provide a sample dataset to help you get started. Download the source media using the following command:
|
||||
|
||||
We provide a sample dataset to help you get started. Download the source media using the following command:
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Image-Vid-Finetune-Src --local_dir=data/Image-Vid-Finetune-Src --repo_type=dataset
|
||||
```
|
||||
|
||||
To preprocess the dataset for fine-tuning or distillation, run:
|
||||
|
||||
```
|
||||
bash scripts/preprocess/preprocess_mochi_data.sh # for mochi
|
||||
bash scripts/preprocess/preprocess_hunyuan_data.sh # for hunyuan
|
||||
@@ -33,16 +33,13 @@ path_to_dataset_folder/
|
||||
Format the JSON file as a list, where each item represents a media source:
|
||||
|
||||
For image media,
|
||||
|
||||
```
|
||||
{
|
||||
"path": "0.jpg",
|
||||
"cap": ["captions"]
|
||||
}
|
||||
```
|
||||
|
||||
For video media,
|
||||
|
||||
For video media,
|
||||
```
|
||||
{
|
||||
"path": "1.mp4",
|
||||
@@ -65,9 +62,7 @@ path_to_media_source_foder,path_to_json_file
|
||||
```
|
||||
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/preprocess_****_data.sh` accordingly and run:
|
||||
|
||||
```
|
||||
bash scripts/preprocess/preprocess_****_data.sh
|
||||
```
|
||||
|
||||
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
|
||||
|
||||
@@ -1,35 +0,0 @@
|
||||
@ECHO OFF
|
||||
|
||||
pushd %~dp0
|
||||
|
||||
REM Command file for Sphinx documentation
|
||||
|
||||
if "%SPHINXBUILD%" == "" (
|
||||
set SPHINXBUILD=sphinx-build
|
||||
)
|
||||
set SOURCEDIR=source
|
||||
set BUILDDIR=build
|
||||
|
||||
%SPHINXBUILD% >NUL 2>NUL
|
||||
if errorlevel 9009 (
|
||||
echo.
|
||||
echo.The 'sphinx-build' command was not found. Make sure you have Sphinx
|
||||
echo.installed, then set the SPHINXBUILD environment variable to point
|
||||
echo.to the full path of the 'sphinx-build' executable. Alternatively you
|
||||
echo.may add the Sphinx directory to PATH.
|
||||
echo.
|
||||
echo.If you don't have Sphinx installed, grab it from
|
||||
echo.https://www.sphinx-doc.org/
|
||||
exit /b 1
|
||||
)
|
||||
|
||||
if "%1" == "" goto help
|
||||
|
||||
%SPHINXBUILD% -M %1 %SOURCEDIR% %BUILDDIR% %SPHINXOPTS% %O%
|
||||
goto end
|
||||
|
||||
:help
|
||||
%SPHINXBUILD% -M help %SOURCEDIR% %BUILDDIR% %SPHINXOPTS% %O%
|
||||
|
||||
:end
|
||||
popd
|
||||
@@ -1,25 +0,0 @@
|
||||
sphinx==6.2.1
|
||||
sphinx-argparse==0.4.0
|
||||
sphinx-book-theme==1.0.1
|
||||
sphinx-copybutton==0.5.2
|
||||
sphinx-design==0.6.1
|
||||
sphinx-togglebutton==0.3.2
|
||||
myst-parser==3.0.1
|
||||
msgspec
|
||||
cloudpickle
|
||||
|
||||
# packages to install to build the documentation
|
||||
cachetools
|
||||
pydantic >= 2.8
|
||||
-f https://download.pytorch.org/whl/cpu
|
||||
torch
|
||||
py-cpuinfo
|
||||
transformers
|
||||
mistral_common >= 1.5.4
|
||||
aiohttp
|
||||
starlette
|
||||
openai # Required by docs/source/serving/openai_compatible_server.md's vllm.entrypoints.openai.cli_args
|
||||
fastapi # Required by docs/source/serving/openai_compatible_server.md's vllm.entrypoints.openai.cli_args
|
||||
partial-json-parser # Required by docs/source/serving/openai_compatible_server.md's vllm.entrypoints.openai.cli_args
|
||||
requests
|
||||
zmq
|
||||
@@ -1,51 +0,0 @@
|
||||
# Seed Parameter Behavior in vLLM
|
||||
|
||||
## Overview
|
||||
|
||||
The `seed` parameter in vLLM is used to control the random states for various random number generators. This parameter can affect the behavior of random operations in user code, especially when working with models in vLLM.
|
||||
|
||||
## Default Behavior
|
||||
|
||||
By default, the `seed` parameter is set to `None`. When the `seed` parameter is `None`, the global random states for `random`, `np.random`, and `torch.manual_seed` are not set. This means that the random operations will behave as expected, without any fixed random states.
|
||||
|
||||
## Specifying a Seed
|
||||
|
||||
If a specific seed value is provided, the global random states for `random`, `np.random`, and `torch.manual_seed` will be set accordingly. This can be useful for reproducibility, as it ensures that the random operations produce the same results across multiple runs.
|
||||
|
||||
## Example Usage
|
||||
|
||||
### Without Specifying a Seed
|
||||
|
||||
```python
|
||||
import random
|
||||
from vllm import LLM
|
||||
|
||||
# Initialize a vLLM model without specifying a seed
|
||||
model = LLM(model="Qwen/Qwen2.5-0.5B-Instruct")
|
||||
|
||||
# Try generating random numbers
|
||||
print(random.randint(0, 100)) # Outputs different numbers across runs
|
||||
```
|
||||
|
||||
### Specifying a Seed
|
||||
|
||||
```python
|
||||
import random
|
||||
from vllm import LLM
|
||||
|
||||
# Initialize a vLLM model with a specific seed
|
||||
model = LLM(model="Qwen/Qwen2.5-0.5B-Instruct", seed=42)
|
||||
|
||||
# Try generating random numbers
|
||||
print(random.randint(0, 100)) # Outputs the same number across runs
|
||||
```
|
||||
|
||||
## Important Notes
|
||||
|
||||
- If the `seed` parameter is not specified, the behavior of global random states remains unaffected.
|
||||
- If a specific seed value is provided, the global random states for `random`, `np.random`, and `torch.manual_seed` will be set to that value.
|
||||
- This behavior can be useful for reproducibility but may lead to non-intuitive behavior if the user is not explicitly aware of it.
|
||||
|
||||
## Conclusion
|
||||
|
||||
Understanding the behavior of the `seed` parameter in vLLM is crucial for ensuring the expected behavior of random operations in your code. By default, the `seed` parameter is set to `None`, which means that the global random states are not affected. However, specifying a seed value can help achieve reproducibility in your experiments.
|
||||
@@ -1,8 +0,0 @@
|
||||
.vertical-table-header th.head:not(.stub) {
|
||||
writing-mode: sideways-lr;
|
||||
white-space: nowrap;
|
||||
max-width: 0;
|
||||
p {
|
||||
margin: 0;
|
||||
}
|
||||
}
|
||||
@@ -1,18 +0,0 @@
|
||||
// Update URL search params when tab is clicked
|
||||
document.addEventListener("DOMContentLoaded", function () {
|
||||
const tabs = document.querySelectorAll(".sd-tab-label");
|
||||
|
||||
function updateURL(tab) {
|
||||
const syncGroup = tab.getAttribute("data-sync-group");
|
||||
const syncId = tab.getAttribute("data-sync-id");
|
||||
if (syncGroup && syncId) {
|
||||
const url = new URL(window.location);
|
||||
url.searchParams.set(syncGroup, syncId);
|
||||
window.history.replaceState(null, "", url);
|
||||
}
|
||||
}
|
||||
|
||||
tabs.forEach(tab => {
|
||||
tab.addEventListener("click", () => updateURL(tab));
|
||||
});
|
||||
});
|
||||
@@ -1,39 +0,0 @@
|
||||
<style>
|
||||
.notification-bar {
|
||||
width: 100vw;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
font-size: 16px;
|
||||
padding: 0 6px 0 6px;
|
||||
}
|
||||
.notification-bar p {
|
||||
margin: 0;
|
||||
}
|
||||
.notification-bar a {
|
||||
font-weight: bold;
|
||||
text-decoration: none;
|
||||
}
|
||||
|
||||
/* Light mode styles (default) */
|
||||
.notification-bar {
|
||||
background-color: #fff3cd;
|
||||
color: #856404;
|
||||
}
|
||||
.notification-bar a {
|
||||
color: #d97706;
|
||||
}
|
||||
|
||||
/* Dark mode styles */
|
||||
html[data-theme=dark] .notification-bar {
|
||||
background-color: #333;
|
||||
color: #ddd;
|
||||
}
|
||||
html[data-theme=dark] .notification-bar a {
|
||||
color: #ffa500; /* Brighter color for visibility */
|
||||
}
|
||||
</style>
|
||||
|
||||
<div class="notification-bar">
|
||||
<p>You are viewing the latest developer preview docs. <a href="https://docs.vllm.ai/en/stable/">Click here</a> to view docs for the latest stable release.</p>
|
||||
</div>
|
||||
@@ -1,260 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Configuration file for the Sphinx documentation builder.
|
||||
#
|
||||
# This file only contains a selection of the most common options. For a full
|
||||
# list see the documentation:
|
||||
# https://www.sphinx-doc.org/en/master/usage/configuration.html
|
||||
|
||||
# -- Path setup --------------------------------------------------------------
|
||||
|
||||
# If extensions (or modules to document with autodoc) are in another directory,
|
||||
# add these directories to sys.path here. If the directory is relative to the
|
||||
# documentation root, use os.path.abspath to make it absolute, like shown here.
|
||||
|
||||
import datetime
|
||||
import inspect
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from typing import Optional
|
||||
|
||||
import requests
|
||||
from sphinx.ext import autodoc
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
sys.path.append(os.path.abspath("../.."))
|
||||
|
||||
# -- Project information -----------------------------------------------------
|
||||
|
||||
project = 'FastVideo'
|
||||
copyright = f'{datetime.datetime.now().year}, FastVideo Team'
|
||||
author = 'the FastVideo Team'
|
||||
|
||||
# -- General configuration ---------------------------------------------------
|
||||
|
||||
# Add any Sphinx extension module names here, as strings. They can be
|
||||
# extensions coming with Sphinx (named 'sphinx.ext.*') or your custom
|
||||
# ones.
|
||||
extensions = [
|
||||
"sphinx.ext.napoleon",
|
||||
"sphinx.ext.linkcode",
|
||||
"sphinx.ext.intersphinx",
|
||||
"sphinx_copybutton",
|
||||
"sphinx.ext.autodoc",
|
||||
"sphinx.ext.autosummary",
|
||||
"myst_parser",
|
||||
"sphinxarg.ext",
|
||||
"sphinx_design",
|
||||
"sphinx_togglebutton",
|
||||
]
|
||||
myst_enable_extensions = [
|
||||
"colon_fence",
|
||||
]
|
||||
|
||||
# Add any paths that contain templates here, relative to this directory.
|
||||
templates_path = ['_templates']
|
||||
|
||||
# List of patterns, relative to source directory, that match files and
|
||||
# directories to ignore when looking for source files.
|
||||
# This pattern also affects html_static_path and html_extra_path.
|
||||
exclude_patterns: list[str] = ["**/*.template.md", "**/*.inc.md"]
|
||||
|
||||
# Exclude the prompt "$" when copying code
|
||||
copybutton_prompt_text = r"\$ "
|
||||
copybutton_prompt_is_regexp = True
|
||||
|
||||
# -- Options for HTML output -------------------------------------------------
|
||||
|
||||
# The theme to use for HTML and HTML Help pages. See the documentation for
|
||||
# a list of builtin themes.
|
||||
#
|
||||
html_title = project
|
||||
html_theme = 'sphinx_book_theme'
|
||||
html_logo = '../../assets/logo.jpg'
|
||||
#html_favicon = 'assets/logos/vllm-logo-only-light.ico'
|
||||
html_theme_options = {
|
||||
'path_to_docs': 'docs/source',
|
||||
'repository_url': 'https://github.com/hao-ai-lab/FastVideo/',
|
||||
'use_repository_button': True,
|
||||
'use_edit_page_button': True,
|
||||
}
|
||||
# Add any paths that contain custom static files (such as style sheets) here,
|
||||
# relative to this directory. They are copied after the builtin static files,
|
||||
# so a file named "default.css" will overwrite the builtin "default.css".
|
||||
html_static_path = ["_static"]
|
||||
html_js_files = ["custom.js"]
|
||||
html_css_files = ["custom.css"]
|
||||
|
||||
myst_url_schemes = {
|
||||
'http': None,
|
||||
'https': None,
|
||||
'mailto': None,
|
||||
'ftp': None,
|
||||
"gh-issue": {
|
||||
"url":
|
||||
"https://github.com/hao-ai-lab/FastVideo/issues/{{path}}#{{fragment}}",
|
||||
"title": "Issue #{{path}}",
|
||||
"classes": ["github"],
|
||||
},
|
||||
"gh-pr": {
|
||||
"url":
|
||||
"https://github.com/hao-ai-lab/FastVideo/pull/{{path}}#{{fragment}}",
|
||||
"title": "Pull Request #{{path}}",
|
||||
"classes": ["github"],
|
||||
},
|
||||
"gh-dir": {
|
||||
"url": "https://github.com/hao-ai-lab/FastVideo/tree/main/{{path}}",
|
||||
"title": "{{path}}",
|
||||
"classes": ["github"],
|
||||
},
|
||||
"gh-file": {
|
||||
"url": "https://github.com/hao-ai-lab/FastVideo/blob/main/{{path}}",
|
||||
"title": "{{path}}",
|
||||
"classes": ["github"],
|
||||
},
|
||||
}
|
||||
|
||||
# see https://docs.readthedocs.io/en/stable/reference/environment-variables.html # noqa
|
||||
READTHEDOCS_VERSION_TYPE = os.environ.get('READTHEDOCS_VERSION_TYPE')
|
||||
if READTHEDOCS_VERSION_TYPE == "tag":
|
||||
# remove the warning banner if the version is a tagged release
|
||||
header_file = os.path.join(os.path.dirname(__file__),
|
||||
"_templates/sections/header.html")
|
||||
# The file might be removed already if the build is triggered multiple times
|
||||
# (readthedocs build both HTML and PDF versions separately)
|
||||
if os.path.exists(header_file):
|
||||
os.remove(header_file)
|
||||
|
||||
|
||||
# Generate additional rst documentation here.
|
||||
def setup(app):
|
||||
from docs.source.generate_examples import generate_examples
|
||||
generate_examples()
|
||||
|
||||
|
||||
_cached_base: str = ""
|
||||
_cached_branch: str = ""
|
||||
|
||||
|
||||
def get_repo_base_and_branch(
|
||||
pr_number: str) -> tuple[Optional[str], Optional[str]]:
|
||||
global _cached_base, _cached_branch
|
||||
if _cached_base and _cached_branch:
|
||||
return _cached_base, _cached_branch
|
||||
|
||||
url = f"https://api.github.com/repos/hao-ai-lab/FastVideo/pulls/{pr_number}"
|
||||
response = requests.get(url)
|
||||
if response.status_code == 200:
|
||||
data = response.json()
|
||||
_cached_base = data['head']['repo']['full_name']
|
||||
_cached_branch = data['head']['ref']
|
||||
return _cached_base, _cached_branch
|
||||
else:
|
||||
logger.error("Failed to fetch PR details: %s", response)
|
||||
return None, None
|
||||
|
||||
|
||||
def linkcode_resolve(domain, info):
|
||||
if domain != 'py':
|
||||
return None
|
||||
if not info['module']:
|
||||
return None
|
||||
module = info['module']
|
||||
|
||||
# try to determine the correct file and line number to link to
|
||||
obj = sys.modules[module]
|
||||
|
||||
# get as specific as we can
|
||||
lineno: int = 0
|
||||
filename: str = ""
|
||||
try:
|
||||
for part in info['fullname'].split('.'):
|
||||
obj = getattr(obj, part)
|
||||
|
||||
if not (inspect.isclass(obj) or inspect.isfunction(obj)
|
||||
or inspect.ismethod(obj)):
|
||||
obj = obj.__class__ # type: ignore[assignment]
|
||||
|
||||
lineno = inspect.getsourcelines(obj)[1]
|
||||
filename = (inspect.getsourcefile(obj)
|
||||
or f"{filename}.py").split("FastVideo/", 1)[1]
|
||||
except Exception:
|
||||
# For some things, like a class member, won't work, so
|
||||
# we'll use the line number of the parent (the class)
|
||||
pass
|
||||
|
||||
if filename.startswith("checkouts/"):
|
||||
# a PR build on readthedocs
|
||||
pr_number = filename.split("/")[1]
|
||||
filename = filename.split("/", 2)[2]
|
||||
base, branch = get_repo_base_and_branch(pr_number)
|
||||
if base and branch:
|
||||
return f"https://github.com/{base}/blob/{branch}/{filename}#L{lineno}"
|
||||
|
||||
# Otherwise, link to the source file on the main branch
|
||||
return f"https://github.com/hao-ai-lab/FastVideo/blob/main/{filename}#L{lineno}"
|
||||
|
||||
|
||||
# Mock out external dependencies here, otherwise the autodoc pages may be blank.
|
||||
autodoc_mock_imports = [
|
||||
"blake3",
|
||||
"compressed_tensors",
|
||||
"cpuinfo",
|
||||
"cv2",
|
||||
"torch",
|
||||
"transformers",
|
||||
"psutil",
|
||||
"prometheus_client",
|
||||
"sentencepiece",
|
||||
"vllm._C",
|
||||
"PIL",
|
||||
"numpy",
|
||||
'triton',
|
||||
"tqdm",
|
||||
"tensorizer",
|
||||
"pynvml",
|
||||
"outlines",
|
||||
"xgrammar",
|
||||
"librosa",
|
||||
"soundfile",
|
||||
"gguf",
|
||||
"lark",
|
||||
"decord",
|
||||
]
|
||||
|
||||
for mock_target in autodoc_mock_imports:
|
||||
if mock_target in sys.modules:
|
||||
logger.info(
|
||||
"Potentially problematic mock target (%s) found; "
|
||||
"autodoc_mock_imports cannot mock modules that have already "
|
||||
"been loaded into sys.modules when the sphinx build starts.",
|
||||
mock_target)
|
||||
|
||||
|
||||
class MockedClassDocumenter(autodoc.ClassDocumenter):
|
||||
"""Remove note about base class when a class is derived from object."""
|
||||
|
||||
def add_line(self, line: str, source: str, *lineno: int) -> None:
|
||||
if line == " Bases: :py:class:`object`":
|
||||
return
|
||||
super().add_line(line, source, *lineno)
|
||||
|
||||
|
||||
autodoc.ClassDocumenter = MockedClassDocumenter
|
||||
|
||||
intersphinx_mapping = {
|
||||
"python": ("https://docs.python.org/3", None),
|
||||
"typing_extensions":
|
||||
("https://typing-extensions.readthedocs.io/en/latest", None),
|
||||
"aiohttp": ("https://docs.aiohttp.org/en/stable", None),
|
||||
"pillow": ("https://pillow.readthedocs.io/en/stable", None),
|
||||
"numpy": ("https://numpy.org/doc/stable", None),
|
||||
"torch": ("https://pytorch.org/docs/stable", None),
|
||||
"psutil": ("https://psutil.readthedocs.io/en/stable", None),
|
||||
}
|
||||
|
||||
autodoc_preserve_defaults = True
|
||||
autodoc_warningiserror = True
|
||||
|
||||
navigation_with_keys = False
|
||||
@@ -1,50 +0,0 @@
|
||||
# Contributing to FastVideo
|
||||
|
||||
Thank you for your interest in contributing to FastVideo. We want to make the process as smooth for you as possible and this is a guide to help get you started!
|
||||
|
||||
Our community is open to everyone and welcomes any contributions no matter how large or small.
|
||||
|
||||
# Developer Environment:
|
||||
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only support Linux and CUDA GPUs, but we hope to support other platforms in the future.
|
||||
|
||||
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
|
||||
|
||||
Install Miniconda:
|
||||
|
||||
```
|
||||
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
|
||||
bash Miniconda3-latest-Linux-x86_64.sh
|
||||
source ~/.bashrc
|
||||
```
|
||||
|
||||
Create and activate a Conda environment for FastVideo:
|
||||
|
||||
```
|
||||
conda create -n fastvideo python=3.10 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
Clone the FastVideo repository and go to the FastVideo directory:
|
||||
|
||||
```
|
||||
git clone https://github.com/vllm-project/vllm.git && cd vllm
|
||||
|
||||
```
|
||||
|
||||
Now you can install FastVideo and setup git hooks for running linting. By using `pre-commit`, the linters will run and have to pass before you'll be able to make a commit.
|
||||
|
||||
```bash
|
||||
pip install -e .[dev]
|
||||
|
||||
# Can also install flash-attn (optional)
|
||||
pip install flash-attn==2.7.0.post2 --no-build-isolation
|
||||
|
||||
# Linting, formatting and static type checking
|
||||
pre-commit install --hook-type pre-commit --hook-type commit-msg
|
||||
|
||||
# You can manually run pre-commit with
|
||||
pre-commit run --all-files
|
||||
|
||||
# Unit tests
|
||||
pytest tests/
|
||||
```
|
||||
@@ -1,246 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import itertools
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
ROOT_DIR = Path(__file__).parent.parent.parent.resolve()
|
||||
ROOT_DIR_RELATIVE = '../../../..'
|
||||
EXAMPLE_DIR = ROOT_DIR / "examples"
|
||||
EXAMPLE_DOC_DIR = ROOT_DIR / "docs/source/getting_started/examples"
|
||||
|
||||
|
||||
def fix_case(text: str) -> str:
|
||||
subs = {
|
||||
"api": "API",
|
||||
"cli": "CLI",
|
||||
"cpu": "CPU",
|
||||
"llm": "LLM",
|
||||
"tpu": "TPU",
|
||||
"aqlm": "AQLM",
|
||||
"gguf": "GGUF",
|
||||
"lora": "LoRA",
|
||||
"rlhf": "RLHF",
|
||||
"vllm": "vLLM",
|
||||
"openai": "OpenAI",
|
||||
"multilora": "MultiLoRA",
|
||||
"mlpspeculator": "MLPSpeculator",
|
||||
r"fp\d+": lambda x: x.group(0).upper(), # e.g. fp16, fp32
|
||||
r"int\d+": lambda x: x.group(0).upper(), # e.g. int8, int16
|
||||
}
|
||||
for pattern, repl in subs.items():
|
||||
text = re.sub(rf'\b{pattern}\b', repl, text,
|
||||
flags=re.IGNORECASE) # type: ignore[call-overload]
|
||||
return text
|
||||
|
||||
|
||||
@dataclass
|
||||
class Index:
|
||||
"""
|
||||
Index class to generate a structured document index.
|
||||
|
||||
Attributes:
|
||||
path (Path): The path save the index file to.
|
||||
title (str): The title of the index.
|
||||
description (str): A brief description of the index.
|
||||
caption (str): An optional caption for the table of contents.
|
||||
maxdepth (int): The maximum depth of the table of contents. Defaults to 1.
|
||||
documents (list[str]): A list of document paths to include in the index. Defaults to an empty list.
|
||||
|
||||
Methods:
|
||||
generate() -> str:
|
||||
Generates the index content as a string in the specified format.
|
||||
""" # noqa: E501
|
||||
path: Path
|
||||
title: str
|
||||
description: str
|
||||
caption: str
|
||||
maxdepth: int = 1
|
||||
documents: list[str] = field(default_factory=list)
|
||||
|
||||
def generate(self) -> str:
|
||||
content = f"# {self.title}\n\n{self.description}\n\n"
|
||||
content += ":::{toctree}\n"
|
||||
content += f":caption: {self.caption}\n:maxdepth: {self.maxdepth}\n"
|
||||
content += "\n".join(self.documents) + "\n:::\n"
|
||||
return content
|
||||
|
||||
|
||||
@dataclass
|
||||
class Example:
|
||||
"""
|
||||
Example class for generating documentation content from a given path.
|
||||
|
||||
Attributes:
|
||||
path (Path): The path to the main directory or file.
|
||||
category (str): The category of the document.
|
||||
main_file (Path): The main file in the directory.
|
||||
other_files (list[Path]): list of other files in the directory.
|
||||
title (str): The title of the document.
|
||||
|
||||
Methods:
|
||||
__post_init__(): Initializes the main_file, other_files, and title attributes.
|
||||
determine_main_file() -> Path: Determines the main file in the given path.
|
||||
determine_other_files() -> list[Path]: Determines other files in the directory excluding the main file.
|
||||
determine_title() -> str: Determines the title of the document.
|
||||
generate() -> str: Generates the documentation content.
|
||||
""" # noqa: E501
|
||||
path: Path
|
||||
category: Optional[str] = None
|
||||
main_file: Path = field(init=False)
|
||||
other_files: list[Path] = field(init=False)
|
||||
title: str = field(init=False)
|
||||
|
||||
def __post_init__(self):
|
||||
self.main_file = self.determine_main_file()
|
||||
self.other_files = self.determine_other_files()
|
||||
self.title = self.determine_title()
|
||||
|
||||
def determine_main_file(self) -> Path:
|
||||
"""
|
||||
Determines the main file in the given path.
|
||||
If the path is a file, it returns the path itself. Otherwise, it searches
|
||||
for Markdown files (*.md) in the directory and returns the first one found.
|
||||
Returns:
|
||||
Path: The main file path, either the original path if it's a file or the first
|
||||
Markdown file found in the directory.
|
||||
Raises:
|
||||
IndexError: If no Markdown files are found in the directory.
|
||||
""" # noqa: E501
|
||||
return self.path if self.path.is_file() else list(
|
||||
self.path.glob("*.md")).pop()
|
||||
|
||||
def determine_other_files(self) -> list[Path]:
|
||||
"""
|
||||
Determine other files in the directory excluding the main file.
|
||||
|
||||
This method checks if the given path is a file. If it is, it returns an empty list.
|
||||
Otherwise, it recursively searches through the directory and returns a list of all
|
||||
files that are not the main file.
|
||||
|
||||
Returns:
|
||||
list[Path]: A list of Path objects representing the other files in the directory.
|
||||
""" # noqa: E501
|
||||
if self.path.is_file():
|
||||
return []
|
||||
is_other_file = lambda file: file.is_file() and file != self.main_file
|
||||
return [file for file in self.path.rglob("*")
|
||||
if is_other_file(file)] # type: ignore[no-untyped-call]
|
||||
|
||||
def determine_title(self) -> str:
|
||||
return fix_case(self.path.stem.replace("_", " ").title())
|
||||
|
||||
def generate(self) -> str:
|
||||
# Convert the path to a relative path from __file__
|
||||
make_relative = lambda path: ROOT_DIR_RELATIVE / path.relative_to(
|
||||
ROOT_DIR)
|
||||
|
||||
content = f"Source <gh-file:{self.path.relative_to(ROOT_DIR)}>.\n\n"
|
||||
include = "include" if self.main_file.suffix == ".md" else \
|
||||
"literalinclude"
|
||||
if include == "literalinclude":
|
||||
content += f"# {self.title}\n\n"
|
||||
content += f":::{{{include}}} {make_relative(self.main_file)}\n" # type: ignore[no-untyped-call]
|
||||
if include == "literalinclude":
|
||||
content += f":language: {self.main_file.suffix[1:]}\n"
|
||||
content += ":::\n\n"
|
||||
|
||||
if not self.other_files:
|
||||
return content
|
||||
|
||||
content += "## Example materials\n\n"
|
||||
for file in sorted(self.other_files):
|
||||
include = "include" if file.suffix == ".md" else "literalinclude"
|
||||
content += f":::{{admonition}} {file.relative_to(self.path)}\n"
|
||||
content += ":class: dropdown\n\n"
|
||||
content += f":::{{{include}}} {make_relative(file)}\n:::\n" # type: ignore[no-untyped-call]
|
||||
content += ":::\n\n"
|
||||
|
||||
return content
|
||||
|
||||
|
||||
def generate_examples():
|
||||
# Create the EXAMPLE_DOC_DIR if it doesn't exist
|
||||
if not EXAMPLE_DOC_DIR.exists():
|
||||
EXAMPLE_DOC_DIR.mkdir(parents=True)
|
||||
|
||||
# Create empty indices
|
||||
examples_index = Index(
|
||||
path=EXAMPLE_DOC_DIR / "examples_index.md",
|
||||
title="Examples",
|
||||
description=
|
||||
"A collection of examples demonstrating usage of FastVideo.\nAll documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.", # noqa: E501
|
||||
caption="Examples",
|
||||
maxdepth=2)
|
||||
# Category indices stored in reverse order because they are inserted into
|
||||
# examples_index.documents at index 0 in order
|
||||
category_indices = {
|
||||
"other":
|
||||
Index(
|
||||
path=EXAMPLE_DOC_DIR / "examples_other_index.md",
|
||||
title="Other",
|
||||
description=
|
||||
"Other examples that don't strongly fit into the online or offline serving categories.", # noqa: E501
|
||||
caption="Examples",
|
||||
),
|
||||
"online_serving":
|
||||
Index(
|
||||
path=EXAMPLE_DOC_DIR / "examples_online_serving_index.md",
|
||||
title="Online Serving",
|
||||
description=
|
||||
"Online serving examples demonstrate how to use FastVideo in an online setting, where the model is queried for predictions in real-time.", # noqa: E501
|
||||
caption="Examples",
|
||||
),
|
||||
"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
|
||||
caption="Examples",
|
||||
),
|
||||
}
|
||||
|
||||
examples = []
|
||||
glob_patterns = ["*.py", "*.md", "*.sh"]
|
||||
# Find categorised examples
|
||||
for category in category_indices:
|
||||
category_dir = EXAMPLE_DIR / category
|
||||
globs = [category_dir.glob(pattern) for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
examples.append(Example(path, category))
|
||||
# Find examples in subdirectories
|
||||
for path in category_dir.glob("*/*.md"):
|
||||
examples.append(Example(path.parent, category))
|
||||
# Find uncategorised examples
|
||||
globs = [EXAMPLE_DIR.glob(pattern) for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
examples.append(Example(path))
|
||||
# Find examples in subdirectories
|
||||
for path in EXAMPLE_DIR.glob("*/*.md"):
|
||||
# Skip categorised examples
|
||||
if path.parent.name in category_indices:
|
||||
continue
|
||||
examples.append(Example(path.parent))
|
||||
|
||||
# Generate the example documentation
|
||||
for example in sorted(examples, key=lambda e: e.path.stem):
|
||||
doc_path = EXAMPLE_DOC_DIR / f"{example.path.stem}.md"
|
||||
with open(doc_path, "w+") as f:
|
||||
f.write(example.generate())
|
||||
# Add the example to the appropriate index
|
||||
assert example.category is not None
|
||||
index = category_indices.get(example.category, examples_index)
|
||||
index.documents.append(example.path.stem)
|
||||
|
||||
# Generate the index files
|
||||
for category_index in category_indices.values():
|
||||
if category_index.documents:
|
||||
examples_index.documents.insert(0, category_index.path.name)
|
||||
with open(category_index.path, "w+") as f:
|
||||
f.write(category_index.generate())
|
||||
|
||||
with open(examples_index.path, "w+") as f:
|
||||
f.write(examples_index.generate())
|
||||
@@ -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 +0,0 @@
|
||||
(fastvideo-installation)=
|
||||
|
||||
# 🔧 Installation
|
||||
The code is tested on Python 3.10.0, CUDA 12.4 and H100.
|
||||
|
||||
```
|
||||
./env_setup.sh fastvideo
|
||||
```
|
||||
|
||||
To try Sliding Tile Attention (optional), please follow the instruction in [here](#sta-installation) to install STA.
|
||||
@@ -1,88 +0,0 @@
|
||||
# Welcome to FastVideo
|
||||
|
||||
:::{figure} ../../assets/logo.jpg
|
||||
:align: center
|
||||
:alt: FastVideo
|
||||
:class: no-scaled-link
|
||||
:width: 60%
|
||||
:::
|
||||
|
||||
:::{raw} html
|
||||
<p style="text-align:center">
|
||||
<strong>FastVideo is a lightweight framework for accelerating large video diffusion models.
|
||||
</strong>
|
||||
</p>
|
||||
|
||||
<p style="text-align:center">
|
||||
<script async defer src="https://buttons.github.io/buttons.js"></script>
|
||||
<a class="github-button" href="https://github.com/hao-ai-lab/FastVideo/" data-show-count="true" data-size="large" aria-label="Star">Star</a>
|
||||
<a class="github-button" href="https://github.com/hao-ai-lab/FastVideo/subscription" data-icon="octicon-eye" data-size="large" aria-label="Watch">Watch</a>
|
||||
<a class="github-button" href="https://github.com/hao-ai-lab/FastVideo/fork" data-icon="octicon-repo-forked" data-size="large" aria-label="Fork">Fork</a>
|
||||
</p>
|
||||
:::
|
||||
|
||||
FastVideo is a lightweight framework for accelerating large video diffusion models developed by the [Hao AI Lab](https://hao-ai-lab.github.io/).
|
||||
|
||||
<div style="text-align: center;">
|
||||
<video controls width="800">
|
||||
<source src="https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1" type="video/mp4">
|
||||
Your browser does not support the video tag.
|
||||
</video>
|
||||
</div>
|
||||
|
||||
FastVideo currently offers: (with more to come)
|
||||
|
||||
- [NEW!] [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- FastHunyuan and FastMochi: consistency distilled video diffusion models for 8x inference speedup.
|
||||
- First open distillation recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
|
||||
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
|
||||
- Scalable training with FSDP, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
- Memory efficient finetuning with LoRA, precomputed latent, and precomputed text embeddings.
|
||||
|
||||
Dev in progress and highly experimental.
|
||||
|
||||
## Documentation
|
||||
|
||||
% How to start using vLLM?
|
||||
|
||||
:::{toctree}
|
||||
:caption: Getting Started
|
||||
:maxdepth: 1
|
||||
|
||||
getting_started/installation
|
||||
getting_started/examples/examples_index
|
||||
:::
|
||||
|
||||
% What is STA Kernel?
|
||||
|
||||
:::{toctree}
|
||||
:caption: Sliding Tile Attention
|
||||
:maxdepth: 1
|
||||
|
||||
sliding_tile_attention/installation
|
||||
sliding_tile_attention/usage
|
||||
sliding_tile_attention/test
|
||||
sliding_tile_attention/demo
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Inference
|
||||
:maxdepth: 1
|
||||
|
||||
inference/stepvideo
|
||||
inference/hunyuanvideo
|
||||
inference/fasthunyuan
|
||||
inference/fastmochi
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Developer Guide
|
||||
:maxdepth: 1
|
||||
|
||||
developer_guide/overview
|
||||
:::
|
||||
|
||||
## Indices and tables
|
||||
|
||||
- {ref}`genindex`
|
||||
- {ref}`modindex`
|
||||
@@ -1,33 +0,0 @@
|
||||
(fasthunyuan)=
|
||||
|
||||
# FastHunyuan
|
||||
## Inference FastHunyuan on single RTX4090
|
||||
We now support NF4 and LLM-INT8 quantized inference using BitsAndBytes for FastHunyuan. With NF4 quantization, inference can be performed on a single RTX 4090 GPU, requiring just 20GB of VRAM.
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan-diffusers --local_dir=data/FastHunyuan-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan_hf_quantization.sh
|
||||
```
|
||||
|
||||
For more information about the VRAM requirements for BitsAndBytes quantization, please refer to the table below (timing measured on an H100 GPU):
|
||||
|
||||
| Configuration | Memory to Init Transformer | Peak Memory After Init Pipeline (Denoise) | Diffusion Time | End-to-End Time |
|
||||
|--------------------------------|----------------------------|--------------------------------------------|----------------|-----------------|
|
||||
| BF16 + Pipeline CPU Offload | 23.883G | 33.744G | 81s | 121.5s |
|
||||
| INT8 + Pipeline CPU Offload | 13.911G | 27.979G | 88s | 116.7s |
|
||||
| NF4 + Pipeline CPU Offload | 9.453G | 19.26G | 78s | 114.5s |
|
||||
|
||||
For improved quality in generated videos, we recommend using a GPU with 80GB of memory to run the BF16 model with the original Hunyuan pipeline. To execute the inference, use the following section:
|
||||
|
||||
## FastHunyuan
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan --local_dir=data/FastHunyuan --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan.sh
|
||||
```
|
||||
|
||||
You can also inference FastHunyuan in the [official Hunyuan github](https://github.com/Tencent/HunyuanVideo).
|
||||
@@ -1,9 +0,0 @@
|
||||
(fastmochi)=
|
||||
|
||||
# FastMochi
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastMochi-diffusers --local_dir=data/FastMochi-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_mochi_sp.sh
|
||||
@@ -1,18 +0,0 @@
|
||||
(hunyuanvideo)=
|
||||
|
||||
# HunyuanVideo
|
||||
## Inference HunyuanVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model
|
||||
```
|
||||
|
||||
We provide two examples in the following script to run inference with STA + [TeaCache](https://github.com/ali-vilab/TeaCache) and STA only.
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_hunyuan_STA.sh
|
||||
```
|
||||
|
||||
## Video Demos using STA + Teacache
|
||||
Visit our [demo website](https://fast-video.github.io/) to explore our complete collection of examples. We shorten a single video generation process from 945s to 317s on H100.
|
||||
@@ -1,16 +0,0 @@
|
||||
(stepvideo)=
|
||||
|
||||
# StepVideo
|
||||
## Inference StepVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
```
|
||||
python scripts/huggingface/download_hf.py --repo_id=stepfun-ai/stepvideo-t2v --local_dir=data/stepvideo-t2v --repo_type=model
|
||||
```
|
||||
|
||||
Use the following scripts to run inference for StepVideo. When using STA for inference, the generated videos will have dimensions of 204×768×768 (currently, this is the only supported shape).
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_stepvideo_STA.sh # Inference stepvideo with STA
|
||||
sh scripts/inference/inference_stepvideo.sh # Inference original stepvideo
|
||||
```
|
||||
@@ -1,11 +0,0 @@
|
||||
(sta-demo)=
|
||||
|
||||
# Demo
|
||||
There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
<div style="text-align: center;">
|
||||
<video controls width="800">
|
||||
<source src="https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747" type="video/mp4">
|
||||
Your browser does not support the video tag.
|
||||
</video>
|
||||
</div>
|
||||
@@ -1,25 +0,0 @@
|
||||
(sta-installation)=
|
||||
|
||||
# Installation
|
||||
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
|
||||
First, install C++20 for ThunderKittens:
|
||||
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
|
||||
Install STA:
|
||||
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
@@ -1,7 +0,0 @@
|
||||
(sta-test)=
|
||||
|
||||
# Test
|
||||
|
||||
```bash
|
||||
python test/test_sta.py
|
||||
```
|
||||
@@ -1,17 +0,0 @@
|
||||
(sta-usage)=
|
||||
|
||||
# Usage
|
||||
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
|
||||
```
|
||||
Executable
+10
@@ -0,0 +1,10 @@
|
||||
#!/bin/bash
|
||||
|
||||
# install torch
|
||||
pip install torch==2.5.0 torchvision --index-url https://download.pytorch.org/whl/cu121
|
||||
|
||||
# install FA2 and diffusers
|
||||
pip install packaging ninja && pip install flash-attn==2.7.0.post2 --no-build-isolation
|
||||
|
||||
# install fastvideo
|
||||
pip install -e .
|
||||
@@ -1,27 +1,24 @@
|
||||
import argparse
|
||||
import torch
|
||||
from accelerate.logging import get_logger
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
from diffusers.utils import export_to_video
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from accelerate.logging import get_logger
|
||||
from diffusers.utils import export_to_video
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from torch.utils.data import DataLoader, Dataset
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.utils.load import load_text_encoder, load_vae
|
||||
|
||||
logger = get_logger(__name__)
|
||||
from torch.utils.data import Dataset
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from torch.utils.data import DataLoader
|
||||
from fastvideo.utils.load import load_text_encoder, load_vae
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
class T5dataset(Dataset):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
json_path,
|
||||
vae_debug,
|
||||
self, json_path, vae_debug,
|
||||
):
|
||||
self.json_path = json_path
|
||||
self.vae_debug = vae_debug
|
||||
@@ -35,7 +32,9 @@ class T5dataset(Dataset):
|
||||
length = self.train_dataset[idx]["length"]
|
||||
if self.vae_debug:
|
||||
latents = torch.load(
|
||||
os.path.join(args.output_dir, "latent", self.train_dataset[idx]["latent_path"]),
|
||||
os.path.join(
|
||||
args.output_dir, "latent", self.train_dataset[idx]["latent_path"]
|
||||
),
|
||||
map_location="cpu",
|
||||
)
|
||||
else:
|
||||
@@ -55,7 +54,9 @@ def main(args):
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
torch.cuda.set_device(local_rank)
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
|
||||
dist.init_process_group(
|
||||
backend="nccl", init_method="env://", world_size=world_size, rank=local_rank
|
||||
)
|
||||
|
||||
videoprocessor = VideoProcessor(vae_scale_factor=8)
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
@@ -69,7 +70,9 @@ def main(args):
|
||||
text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
|
||||
vae.enable_tiling()
|
||||
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
|
||||
sampler = DistributedSampler(
|
||||
train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True
|
||||
)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
@@ -81,16 +84,23 @@ def main(args):
|
||||
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
|
||||
with torch.inference_mode():
|
||||
with torch.autocast("cuda", dtype=autocast_type):
|
||||
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(prompt=data["caption"], )
|
||||
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(
|
||||
prompt=data["caption"],
|
||||
)
|
||||
if args.vae_debug:
|
||||
latents = data["latents"]
|
||||
video = vae.decode(latents.to(device), return_dict=False)[0]
|
||||
video = videoprocessor.postprocess_video(video)
|
||||
for idx, video_name in enumerate(data["filename"]):
|
||||
prompt_embed_path = os.path.join(args.output_dir, "prompt_embed", video_name + ".pt")
|
||||
video_path = os.path.join(args.output_dir, "video", video_name + ".mp4")
|
||||
prompt_attention_mask_path = os.path.join(args.output_dir, "prompt_attention_mask",
|
||||
video_name + ".pt")
|
||||
prompt_embed_path = os.path.join(
|
||||
args.output_dir, "prompt_embed", video_name + ".pt"
|
||||
)
|
||||
video_path = os.path.join(
|
||||
args.output_dir, "video", video_name + ".mp4"
|
||||
)
|
||||
prompt_attention_mask_path = os.path.join(
|
||||
args.output_dir, "prompt_attention_mask", video_name + ".pt"
|
||||
)
|
||||
# save latent
|
||||
torch.save(prompt_embeds[idx], prompt_embed_path)
|
||||
torch.save(prompt_attention_mask[idx], prompt_attention_mask_path)
|
||||
|
||||
@@ -1,16 +1,18 @@
|
||||
from fastvideo.dataset import getdataset
|
||||
from torch.utils.data import DataLoader
|
||||
from fastvideo.utils.dataset_utils import Collate
|
||||
import argparse
|
||||
import torch
|
||||
from accelerate import Accelerator
|
||||
from accelerate.logging import get_logger
|
||||
from accelerate.utils import ProjectConfiguration
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
from diffusers import AutoencoderKLMochi
|
||||
import torch.distributed as dist
|
||||
from accelerate.logging import get_logger
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.dataset import getdataset
|
||||
from fastvideo.utils.load import load_vae
|
||||
from tqdm import tqdm
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
@@ -20,7 +22,9 @@ def main(args):
|
||||
world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
print("world_size", world_size, "local rank", local_rank)
|
||||
train_dataset = getdataset(args)
|
||||
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
|
||||
sampler = DistributedSampler(
|
||||
train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True
|
||||
)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
@@ -28,10 +32,12 @@ def main(args):
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
|
||||
encoder_device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
encoder_device = torch.device(f"cuda" if torch.cuda.is_available() else "cpu")
|
||||
torch.cuda.set_device(local_rank)
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
|
||||
dist.init_process_group(
|
||||
backend="nccl", init_method="env://", world_size=world_size, rank=local_rank
|
||||
)
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
|
||||
vae.enable_tiling()
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
@@ -41,10 +47,14 @@ def main(args):
|
||||
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
|
||||
with torch.inference_mode():
|
||||
with torch.autocast("cuda", dtype=autocast_type):
|
||||
latents = vae.encode(data["pixel_values"].to(encoder_device))["latent_dist"].sample()
|
||||
latents = vae.encode(data["pixel_values"].to(encoder_device))[
|
||||
"latent_dist"
|
||||
].sample()
|
||||
for idx, video_path in enumerate(data["path"]):
|
||||
video_name = os.path.basename(video_path).split(".")[0]
|
||||
latent_path = os.path.join(args.output_dir, "latent", video_name + ".pt")
|
||||
latent_path = os.path.join(
|
||||
args.output_dir, "latent", video_name + ".pt"
|
||||
)
|
||||
torch.save(latents[idx].to(torch.bfloat16), latent_path)
|
||||
item = {}
|
||||
item["length"] = latents[idx].shape[1]
|
||||
@@ -81,7 +91,9 @@ if __name__ == "__main__":
|
||||
default=16,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
|
||||
parser.add_argument(
|
||||
"--num_latent_t", type=int, default=28, help="Number of latent timesteps."
|
||||
)
|
||||
parser.add_argument("--max_height", type=int, default=480)
|
||||
parser.add_argument("--max_width", type=int, default=848)
|
||||
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
|
||||
@@ -107,8 +119,10 @@ if __name__ == "__main__":
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
help=(
|
||||
"[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
|
||||
),
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
@@ -1,13 +1,19 @@
|
||||
import argparse
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from accelerate.logging import get_logger
|
||||
|
||||
from fastvideo.utils.load import load_text_encoder
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
from diffusers.utils import export_to_video
|
||||
import json
|
||||
import os
|
||||
import torch.distributed as dist
|
||||
|
||||
logger = get_logger(__name__)
|
||||
from torch.utils.data import Dataset
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from torch.utils.data import DataLoader
|
||||
from fastvideo.utils.load import load_text_encoder, load_vae
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def main(args):
|
||||
@@ -18,7 +24,9 @@ def main(args):
|
||||
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
|
||||
torch.cuda.set_device(local_rank)
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
|
||||
dist.init_process_group(
|
||||
backend="nccl", init_method="env://", world_size=world_size, rank=local_rank
|
||||
)
|
||||
|
||||
text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
|
||||
autocast_type = torch.float16 if args.model_type == "hunyuan" else torch.bfloat16
|
||||
@@ -29,17 +37,23 @@ def main(args):
|
||||
os.path.join(args.output_dir, "validation", "prompt_attention_mask"),
|
||||
exist_ok=True,
|
||||
)
|
||||
os.makedirs(os.path.join(args.output_dir, "validation", "prompt_embed"), exist_ok=True)
|
||||
|
||||
os.makedirs(
|
||||
os.path.join(args.output_dir, "validation", "prompt_embed"), exist_ok=True
|
||||
)
|
||||
json_data = []
|
||||
with open(args.validation_prompt_txt, "r", encoding="utf-8") as file:
|
||||
lines = file.readlines()
|
||||
prompts = [line.strip() for line in lines]
|
||||
for prompt in prompts:
|
||||
with torch.inference_mode():
|
||||
with torch.autocast("cuda", dtype=autocast_type):
|
||||
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(prompt)
|
||||
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(
|
||||
prompt
|
||||
)
|
||||
file_name = prompt.split(".")[0]
|
||||
prompt_embed_path = os.path.join(args.output_dir, "validation", "prompt_embed", f"{file_name}.pt")
|
||||
prompt_embed_path = os.path.join(
|
||||
args.output_dir, "validation", "prompt_embed", f"{file_name}.pt"
|
||||
)
|
||||
prompt_attention_mask_path = os.path.join(
|
||||
args.output_dir,
|
||||
"validation",
|
||||
|
||||
@@ -1,9 +1,14 @@
|
||||
from torchvision import transforms
|
||||
from torchvision.transforms import Lambda
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from torchvision import transforms
|
||||
from torchvision.transforms import Lambda
|
||||
from fastvideo.dataset.t2v_datasets import T2V_dataset
|
||||
from fastvideo.dataset.transform import CenterCropResizeVideo, Normalize255, TemporalRandomCrop
|
||||
from fastvideo.dataset.latent_datasets import LatentDataset
|
||||
from fastvideo.dataset.transform import (
|
||||
Normalize255,
|
||||
TemporalRandomCrop,
|
||||
CenterCropResizeVideo,
|
||||
)
|
||||
|
||||
|
||||
def getdataset(args):
|
||||
@@ -15,17 +20,26 @@ def getdataset(args):
|
||||
resize = [
|
||||
CenterCropResizeVideo((args.max_height, args.max_width)),
|
||||
]
|
||||
transform = transforms.Compose([
|
||||
# Normalize255(),
|
||||
*resize,
|
||||
])
|
||||
transform_topcrop = transforms.Compose([
|
||||
Normalize255(),
|
||||
*resize_topcrop,
|
||||
norm_fun,
|
||||
])
|
||||
transform = transforms.Compose(
|
||||
[
|
||||
# Normalize255(),
|
||||
*resize,
|
||||
# RandomHorizontalFlipVideo(p=0.5), # in case their caption have position decription
|
||||
# norm_fun
|
||||
]
|
||||
)
|
||||
transform_topcrop = transforms.Compose(
|
||||
[
|
||||
Normalize255(),
|
||||
*resize_topcrop,
|
||||
# RandomHorizontalFlipVideo(p=0.5), # in case their caption have position decription
|
||||
norm_fun,
|
||||
]
|
||||
)
|
||||
# tokenizer = AutoTokenizer.from_pretrained("/storage/ongoing/new/Open-Sora-Plan/cache_dir/mt5-xxl", cache_dir=args.cache_dir)
|
||||
tokenizer = AutoTokenizer.from_pretrained(args.text_encoder_name, cache_dir=args.cache_dir)
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
args.text_encoder_name, cache_dir=args.cache_dir
|
||||
)
|
||||
if args.dataset == "t2v":
|
||||
return T2V_dataset(
|
||||
args,
|
||||
@@ -39,12 +53,10 @@ def getdataset(args):
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import random
|
||||
|
||||
from accelerate import Accelerator
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.dataset.t2v_datasets import dataset_prog
|
||||
import random
|
||||
from tqdm import tqdm
|
||||
|
||||
args = type(
|
||||
"args",
|
||||
@@ -80,7 +92,9 @@ if __name__ == "__main__":
|
||||
zero = 0
|
||||
for idx in tqdm(range(num)):
|
||||
image_data = dataset_prog.img_cap_list[idx]
|
||||
caps = [i["cap"] if isinstance(i["cap"], list) else [i["cap"]] for i in image_data]
|
||||
caps = [
|
||||
i["cap"] if isinstance(i["cap"], list) else [i["cap"]] for i in image_data
|
||||
]
|
||||
try:
|
||||
caps = [[random.choice(i)] for i in caps]
|
||||
except Exception as e:
|
||||
|
||||
@@ -1,18 +1,13 @@
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
|
||||
import torch
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
|
||||
class LatentDataset(Dataset):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
json_path,
|
||||
num_latent_t,
|
||||
cfg_rate,
|
||||
self, json_path, num_latent_t, cfg_rate,
|
||||
):
|
||||
# data_merge_path: video_dir, latent_dir, prompt_embed_dir, json_path
|
||||
self.json_path = json_path
|
||||
@@ -21,7 +16,9 @@ class LatentDataset(Dataset):
|
||||
self.video_dir = os.path.join(self.datase_dir_path, "video")
|
||||
self.latent_dir = os.path.join(self.datase_dir_path, "latent")
|
||||
self.prompt_embed_dir = os.path.join(self.datase_dir_path, "prompt_embed")
|
||||
self.prompt_attention_mask_dir = os.path.join(self.datase_dir_path, "prompt_attention_mask")
|
||||
self.prompt_attention_mask_dir = os.path.join(
|
||||
self.datase_dir_path, "prompt_attention_mask"
|
||||
)
|
||||
with open(self.json_path, "r") as f:
|
||||
self.data_anno = json.load(f)
|
||||
# json.load(f) already keeps the order
|
||||
@@ -31,7 +28,10 @@ class LatentDataset(Dataset):
|
||||
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
|
||||
# 256 zeros
|
||||
self.uncond_prompt_mask = torch.zeros(256).bool()
|
||||
self.lengths = [data_item["length"] if "length" in data_item else 1 for data_item in self.data_anno]
|
||||
self.lengths = [
|
||||
data_item["length"] if "length" in data_item else 1
|
||||
for data_item in self.data_anno
|
||||
]
|
||||
|
||||
def __getitem__(self, idx):
|
||||
latent_file = self.data_anno[idx]["latent_path"]
|
||||
@@ -43,7 +43,7 @@ class LatentDataset(Dataset):
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
latent = latent.squeeze(0)[:, -self.num_latent_t:]
|
||||
latent = latent.squeeze(0)[:, -self.num_latent_t :]
|
||||
if random.random() < self.cfg_rate:
|
||||
prompt_embed = self.uncond_prompt_embed
|
||||
prompt_attention_mask = self.uncond_prompt_mask
|
||||
@@ -54,7 +54,9 @@ class LatentDataset(Dataset):
|
||||
weights_only=True,
|
||||
)
|
||||
prompt_attention_mask = torch.load(
|
||||
os.path.join(self.prompt_attention_mask_dir, prompt_attention_mask_file),
|
||||
os.path.join(
|
||||
self.prompt_attention_mask_dir, prompt_attention_mask_file
|
||||
),
|
||||
map_location="cpu",
|
||||
weights_only=True,
|
||||
)
|
||||
@@ -87,15 +89,16 @@ def latent_collate_function(batch):
|
||||
0,
|
||||
max_w - latent.shape[3],
|
||||
),
|
||||
) for latent in latents
|
||||
)
|
||||
for latent in latents
|
||||
]
|
||||
# attn mask
|
||||
latent_attn_mask = torch.ones(len(latents), max_t, max_h, max_w)
|
||||
# set to 0 if padding
|
||||
for i, latent in enumerate(latents):
|
||||
latent_attn_mask[i, latent.shape[1]:, :, :] = 0
|
||||
latent_attn_mask[i, :, latent.shape[2]:, :] = 0
|
||||
latent_attn_mask[i, :, :, latent.shape[3]:] = 0
|
||||
latent_attn_mask[i, latent.shape[1] :, :, :] = 0
|
||||
latent_attn_mask[i, :, latent.shape[2] :, :] = 0
|
||||
latent_attn_mask[i, :, :, latent.shape[3] :] = 0
|
||||
|
||||
prompt_embeds = torch.stack(prompt_embeds, dim=0)
|
||||
prompt_attention_masks = torch.stack(prompt_attention_masks, dim=0)
|
||||
@@ -104,8 +107,10 @@ def latent_collate_function(batch):
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt", num_latent_t=28)
|
||||
dataloader = torch.utils.data.DataLoader(dataset, batch_size=2, shuffle=False, collate_fn=latent_collate_function)
|
||||
dataset = LatentDataset("data/HD-Mixkit-Finetune-Hunyuan/videos2caption.json", num_latent_t=8, cfg_rate=6)
|
||||
dataloader = torch.utils.data.DataLoader(
|
||||
dataset, batch_size=2, shuffle=False, collate_fn=latent_collate_function
|
||||
)
|
||||
for latent, prompt_embed, latent_attn_mask, prompt_attention_mask in dataloader:
|
||||
print(
|
||||
latent.shape,
|
||||
|
||||
@@ -1,18 +1,18 @@
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from collections import Counter
|
||||
from os.path import join as opj
|
||||
|
||||
import os, io, csv, math, random
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
from torch.utils.data import Dataset
|
||||
from decord import VideoReader
|
||||
from os.path import join as opj
|
||||
from collections import Counter
|
||||
|
||||
import torch
|
||||
from torch.utils.data.dataset import Dataset
|
||||
from torch.utils.data import DataLoader, Dataset, get_worker_info
|
||||
from tqdm import tqdm
|
||||
from PIL import Image
|
||||
from fastvideo.utils.dataset_utils import DecordInit
|
||||
import torchvision
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
|
||||
|
||||
@@ -27,7 +27,6 @@ class SingletonMeta(type):
|
||||
|
||||
|
||||
class DataSetProg(metaclass=SingletonMeta):
|
||||
|
||||
def __init__(self):
|
||||
self.cap_list = []
|
||||
self.elements = []
|
||||
@@ -57,7 +56,9 @@ class DataSetProg(metaclass=SingletonMeta):
|
||||
else:
|
||||
worker_id = work_info.id
|
||||
|
||||
idx = self.worker_elements[worker_id][self.n_used_elements[worker_id] % len(self.worker_elements[worker_id])]
|
||||
idx = self.worker_elements[worker_id][
|
||||
self.n_used_elements[worker_id] % len(self.worker_elements[worker_id])
|
||||
]
|
||||
self.n_used_elements[worker_id] += 1
|
||||
return idx
|
||||
|
||||
@@ -72,7 +73,6 @@ def filter_resolution(h, w, max_h_div_w_ratio=17 / 16, min_h_div_w_ratio=8 / 16)
|
||||
|
||||
|
||||
class T2V_dataset(Dataset):
|
||||
|
||||
def __init__(self, args, transform, temporal_sample, tokenizer, transform_topcrop):
|
||||
self.data = args.data_merge_path
|
||||
self.num_frames = args.num_frames
|
||||
@@ -92,7 +92,7 @@ class T2V_dataset(Dataset):
|
||||
self.v_decoder = DecordInit()
|
||||
self.video_length_tolerance_range = args.video_length_tolerance_range
|
||||
self.support_Chinese = True
|
||||
if "mt5" not in args.text_encoder_name:
|
||||
if not ("mt5" in args.text_encoder_name):
|
||||
self.support_Chinese = False
|
||||
|
||||
cap_list = self.get_cap_list()
|
||||
@@ -129,7 +129,9 @@ class T2V_dataset(Dataset):
|
||||
video_path = dataset_prog.cap_list[idx]["path"]
|
||||
assert os.path.exists(video_path), f"file {video_path} do not exist!"
|
||||
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
|
||||
torchvision_video, _, metadata = torchvision.io.read_video(video_path, output_format="TCHW")
|
||||
torchvision_video, _, metadata = torchvision.io.read_video(
|
||||
video_path, output_format="TCHW"
|
||||
)
|
||||
video = torchvision_video[frame_indices]
|
||||
video = self.transform(video)
|
||||
video = rearrange(video, "t c h w -> c t h w")
|
||||
@@ -178,13 +180,20 @@ class T2V_dataset(Dataset):
|
||||
# h, w = i.shape[-2:]
|
||||
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
|
||||
|
||||
image = (self.transform_topcrop(image) if "human_images" in image_data["path"] else self.transform(image)
|
||||
) # [1 C H W] -> num_img [1 C H W]
|
||||
image = (
|
||||
self.transform_topcrop(image)
|
||||
if "human_images" in image_data["path"]
|
||||
else self.transform(image)
|
||||
) # [1 C H W] -> num_img [1 C H W]
|
||||
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
|
||||
|
||||
image = image.float() / 127.5 - 1.0
|
||||
|
||||
caps = (image_data["cap"] if isinstance(image_data["cap"], list) else [image_data["cap"]])
|
||||
caps = (
|
||||
image_data["cap"]
|
||||
if isinstance(image_data["cap"], list)
|
||||
else [image_data["cap"]]
|
||||
)
|
||||
caps = [random.choice(caps)]
|
||||
text = caps
|
||||
input_ids, cond_mask = [], []
|
||||
@@ -238,7 +247,10 @@ class T2V_dataset(Dataset):
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
else:
|
||||
if (resolution.get("height", None) is None or resolution.get("width", None) is None):
|
||||
if (
|
||||
resolution.get("height", None) is None
|
||||
or resolution.get("width", None) is None
|
||||
):
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
height, width = i["resolution"]["height"], i["resolution"]["width"]
|
||||
@@ -259,18 +271,23 @@ class T2V_dataset(Dataset):
|
||||
i["num_frames"] = math.ceil(fps * duration)
|
||||
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
|
||||
if i["num_frames"] / fps > self.video_length_tolerance_range * (
|
||||
self.num_frames / self.train_fps *
|
||||
self.speed_factor): # too long video is not suitable for this training stage (self.num_frames)
|
||||
self.num_frames / self.train_fps * self.speed_factor
|
||||
): # too long video is not suitable for this training stage (self.num_frames)
|
||||
cnt_too_long += 1
|
||||
continue
|
||||
|
||||
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
|
||||
frame_interval = fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, i["num_frames"], frame_interval).astype(int)
|
||||
frame_indices = np.arange(
|
||||
start_frame_idx, i["num_frames"], frame_interval
|
||||
).astype(int)
|
||||
|
||||
# comment out it to enable dynamic frames training
|
||||
if (len(frame_indices) < self.num_frames and random.random() < self.drop_short_ratio):
|
||||
if (
|
||||
len(frame_indices) < self.num_frames
|
||||
and random.random() < self.drop_short_ratio
|
||||
):
|
||||
cnt_too_short += 1
|
||||
continue
|
||||
|
||||
@@ -281,7 +298,9 @@ class T2V_dataset(Dataset):
|
||||
# frame_indices = frame_indices[:self.num_frames] # head crop
|
||||
i["sample_frame_index"] = frame_indices.tolist()
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = len(i["sample_frame_index"]) # will use in dataloader(group sampler)
|
||||
i["sample_num_frames"] = len(
|
||||
i["sample_frame_index"]
|
||||
) # will use in dataloader(group sampler)
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
elif path.endswith(".jpg"): # image
|
||||
cnt_img += 1
|
||||
@@ -290,13 +309,15 @@ class T2V_dataset(Dataset):
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
else:
|
||||
raise NameError(
|
||||
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image")
|
||||
f"Unknown file extention {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
|
||||
)
|
||||
# import ipdb;ipdb.set_trace()
|
||||
main_print(
|
||||
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
|
||||
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
|
||||
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
|
||||
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}")
|
||||
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}"
|
||||
)
|
||||
return new_cap_list, sample_num_frames
|
||||
|
||||
def decord_read(self, path, frame_indices):
|
||||
@@ -309,7 +330,9 @@ class T2V_dataset(Dataset):
|
||||
def read_jsons(self, data):
|
||||
cap_lists = []
|
||||
with open(data, "r") as f:
|
||||
folder_anno = [i.strip().split(",") for i in f.readlines() if len(i.strip()) > 0]
|
||||
folder_anno = [
|
||||
i.strip().split(",") for i in f.readlines() if len(i.strip()) > 0
|
||||
]
|
||||
print(folder_anno)
|
||||
for folder, anno in folder_anno:
|
||||
with open(anno, "r") as f:
|
||||
|
||||
@@ -1,8 +1,7 @@
|
||||
import numbers
|
||||
import random
|
||||
|
||||
import torch
|
||||
from PIL import Image
|
||||
import random
|
||||
import numbers
|
||||
from torchvision.transforms import RandomCrop, RandomResizedCrop
|
||||
|
||||
|
||||
def _is_tensor_video_clip(clip):
|
||||
@@ -21,15 +20,21 @@ def center_crop_arr(pil_image, image_size):
|
||||
https://github.com/openai/guided-diffusion/blob/8fb3ad9197f16bbc40620447b2742e13458d2831/guided_diffusion/image_datasets.py#L126
|
||||
"""
|
||||
while min(*pil_image.size) >= 2 * image_size:
|
||||
pil_image = pil_image.resize(tuple(x // 2 for x in pil_image.size), resample=Image.BOX)
|
||||
pil_image = pil_image.resize(
|
||||
tuple(x // 2 for x in pil_image.size), resample=Image.BOX
|
||||
)
|
||||
|
||||
scale = image_size / min(*pil_image.size)
|
||||
pil_image = pil_image.resize(tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC)
|
||||
pil_image = pil_image.resize(
|
||||
tuple(round(x * scale) for x in pil_image.size), resample=Image.BICUBIC
|
||||
)
|
||||
|
||||
arr = np.array(pil_image)
|
||||
crop_y = (arr.shape[0] - image_size) // 2
|
||||
crop_x = (arr.shape[1] - image_size) // 2
|
||||
return Image.fromarray(arr[crop_y:crop_y + image_size, crop_x:crop_x + image_size])
|
||||
return Image.fromarray(
|
||||
arr[crop_y : crop_y + image_size, crop_x : crop_x + image_size]
|
||||
)
|
||||
|
||||
|
||||
def crop(clip, i, j, h, w):
|
||||
@@ -39,12 +44,14 @@ def crop(clip, i, j, h, w):
|
||||
"""
|
||||
if len(clip.size()) != 4:
|
||||
raise ValueError("clip should be a 4D tensor")
|
||||
return clip[..., i:i + h, j:j + w]
|
||||
return clip[..., i : i + h, j : j + w]
|
||||
|
||||
|
||||
def resize(clip, target_size, interpolation_mode):
|
||||
if len(target_size) != 2:
|
||||
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
|
||||
raise ValueError(
|
||||
f"target size should be tuple (height, width), instead got {target_size}"
|
||||
)
|
||||
return torch.nn.functional.interpolate(
|
||||
clip,
|
||||
size=target_size,
|
||||
@@ -56,7 +63,9 @@ def resize(clip, target_size, interpolation_mode):
|
||||
|
||||
def resize_scale(clip, target_size, interpolation_mode):
|
||||
if len(target_size) != 2:
|
||||
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
|
||||
raise ValueError(
|
||||
f"target size should be tuple (height, width), instead got {target_size}"
|
||||
)
|
||||
H, W = clip.size(-2), clip.size(-1)
|
||||
scale_ = target_size[0] / min(H, W)
|
||||
return torch.nn.functional.interpolate(
|
||||
@@ -144,14 +153,16 @@ def random_shift_crop(clip):
|
||||
h, w = clip.size(-2), clip.size(-1)
|
||||
|
||||
if h <= w:
|
||||
long_edge = w
|
||||
short_edge = h
|
||||
else:
|
||||
long_edge = h
|
||||
short_edge = w
|
||||
|
||||
th, tw = short_edge, short_edge
|
||||
|
||||
i = torch.randint(0, h - th + 1, size=(1, )).item()
|
||||
j = torch.randint(0, w - tw + 1, size=(1, )).item()
|
||||
i = torch.randint(0, h - th + 1, size=(1,)).item()
|
||||
j = torch.randint(0, w - tw + 1, size=(1,)).item()
|
||||
return crop(clip, i, j, th, tw)
|
||||
|
||||
|
||||
@@ -166,7 +177,9 @@ def normalize_video(clip):
|
||||
"""
|
||||
_is_tensor_video_clip(clip)
|
||||
if not clip.dtype == torch.uint8:
|
||||
raise TypeError("clip tensor should have data type uint8. Got %s" % str(clip.dtype))
|
||||
raise TypeError(
|
||||
"clip tensor should have data type uint8. Got %s" % str(clip.dtype)
|
||||
)
|
||||
# return clip.float().permute(3, 0, 1, 2) / 255.0
|
||||
return clip.float() / 255.0
|
||||
|
||||
@@ -204,7 +217,6 @@ def hflip(clip):
|
||||
|
||||
|
||||
class RandomCropVideo:
|
||||
|
||||
def __init__(self, size):
|
||||
if isinstance(size, numbers.Number):
|
||||
self.size = (int(size), int(size))
|
||||
@@ -227,13 +239,15 @@ class RandomCropVideo:
|
||||
th, tw = self.size
|
||||
|
||||
if h < th or w < tw:
|
||||
raise ValueError(f"Required crop size {(th, tw)} is larger than input image size {(h, w)}")
|
||||
raise ValueError(
|
||||
f"Required crop size {(th, tw)} is larger than input image size {(h, w)}"
|
||||
)
|
||||
|
||||
if w == tw and h == th:
|
||||
return 0, 0, h, w
|
||||
|
||||
i = torch.randint(0, h - th + 1, size=(1, )).item()
|
||||
j = torch.randint(0, w - tw + 1, size=(1, )).item()
|
||||
i = torch.randint(0, h - th + 1, size=(1,)).item()
|
||||
j = torch.randint(0, w - tw + 1, size=(1,)).item()
|
||||
|
||||
return i, j, th, tw
|
||||
|
||||
@@ -242,7 +256,6 @@ class RandomCropVideo:
|
||||
|
||||
|
||||
class SpatialStrideCropVideo:
|
||||
|
||||
def __init__(self, stride):
|
||||
self.stride = stride
|
||||
|
||||
@@ -275,10 +288,7 @@ class LongSideResizeVideo:
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
skip_low_resolution=False,
|
||||
interpolation_mode="bilinear",
|
||||
self, size, skip_low_resolution=False, interpolation_mode="bilinear",
|
||||
):
|
||||
self.size = size
|
||||
self.skip_low_resolution = skip_low_resolution
|
||||
@@ -301,7 +311,9 @@ class LongSideResizeVideo:
|
||||
else:
|
||||
h = int(h * self.size / w)
|
||||
w = self.size
|
||||
resize_clip = resize(clip, target_size=(h, w), interpolation_mode=self.interpolation_mode)
|
||||
resize_clip = resize(
|
||||
clip, target_size=(h, w), interpolation_mode=self.interpolation_mode
|
||||
)
|
||||
return resize_clip
|
||||
|
||||
def __repr__(self) -> str:
|
||||
@@ -315,13 +327,12 @@ class CenterCropResizeVideo:
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
top_crop=False,
|
||||
interpolation_mode="bilinear",
|
||||
self, size, top_crop=False, interpolation_mode="bilinear",
|
||||
):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
raise ValueError(
|
||||
f"size should be tuple (height, width), instead got {size}"
|
||||
)
|
||||
self.size = size
|
||||
self.top_crop = top_crop
|
||||
self.interpolation_mode = interpolation_mode
|
||||
@@ -335,7 +346,9 @@ class CenterCropResizeVideo:
|
||||
size is (T, C, crop_size, crop_size)
|
||||
"""
|
||||
# clip_center_crop = center_crop_using_short_edge(clip)
|
||||
clip_center_crop = center_crop_th_tw(clip, self.size[0], self.size[1], top_crop=self.top_crop)
|
||||
clip_center_crop = center_crop_th_tw(
|
||||
clip, self.size[0], self.size[1], top_crop=self.top_crop
|
||||
)
|
||||
# import ipdb;ipdb.set_trace()
|
||||
clip_center_crop_resize = resize(
|
||||
clip_center_crop,
|
||||
@@ -355,13 +368,13 @@ class UCFCenterCropVideo:
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
self, size, interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
raise ValueError(
|
||||
f"size should be tuple (height, width), instead got {size}"
|
||||
)
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
@@ -376,7 +389,9 @@ class UCFCenterCropVideo:
|
||||
torch.tensor: scale resized / center cropped video clip.
|
||||
size is (T, C, crop_size, crop_size)
|
||||
"""
|
||||
clip_resize = resize_scale(clip=clip, target_size=self.size, interpolation_mode=self.interpolation_mode)
|
||||
clip_resize = resize_scale(
|
||||
clip=clip, target_size=self.size, interpolation_mode=self.interpolation_mode
|
||||
)
|
||||
clip_center_crop = center_crop(clip_resize, self.size)
|
||||
return clip_center_crop
|
||||
|
||||
@@ -390,13 +405,13 @@ class KineticsRandomCropResizeVideo:
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
self, size, interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
raise ValueError(
|
||||
f"size should be tuple (height, width), instead got {size}"
|
||||
)
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
@@ -410,15 +425,14 @@ class KineticsRandomCropResizeVideo:
|
||||
|
||||
|
||||
class CenterCropVideo:
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
size,
|
||||
interpolation_mode="bilinear",
|
||||
self, size, interpolation_mode="bilinear",
|
||||
):
|
||||
if isinstance(size, tuple):
|
||||
if len(size) != 2:
|
||||
raise ValueError(f"size should be tuple (height, width), instead got {size}")
|
||||
raise ValueError(
|
||||
f"size should be tuple (height, width), instead got {size}"
|
||||
)
|
||||
self.size = size
|
||||
else:
|
||||
self.size = (size, size)
|
||||
@@ -545,7 +559,9 @@ class DynamicSampleDuration(object):
|
||||
def __call__(self, t, h, w):
|
||||
if self.extra_1:
|
||||
t = t - 1
|
||||
truncate_t_list = list(range(t + 1))[t // 2:][::self.t_stride] # need half at least
|
||||
truncate_t_list = list(range(t + 1))[t // 2 :][
|
||||
:: self.t_stride
|
||||
] # need half at least
|
||||
truncate_t = random.choice(truncate_t_list)
|
||||
if self.extra_1:
|
||||
truncate_t = truncate_t + 1
|
||||
@@ -553,22 +569,27 @@ class DynamicSampleDuration(object):
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
from torchvision import transforms
|
||||
import torchvision.io as io
|
||||
import numpy as np
|
||||
from torchvision.utils import save_image
|
||||
import os
|
||||
|
||||
import numpy as np
|
||||
import torchvision.io as io
|
||||
from torchvision import transforms
|
||||
from torchvision.utils import save_image
|
||||
vframes, aframes, info = io.read_video(
|
||||
filename="./v_Archery_g01_c03.avi", pts_unit="sec", output_format="TCHW"
|
||||
)
|
||||
|
||||
vframes, aframes, info = io.read_video(filename="./v_Archery_g01_c03.avi", pts_unit="sec", output_format="TCHW")
|
||||
|
||||
trans = transforms.Compose([
|
||||
Normalize255(),
|
||||
RandomHorizontalFlipVideo(),
|
||||
UCFCenterCropVideo(512),
|
||||
# NormalizeVideo(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
|
||||
])
|
||||
trans = transforms.Compose(
|
||||
[
|
||||
Normalize255(),
|
||||
RandomHorizontalFlipVideo(),
|
||||
UCFCenterCropVideo(512),
|
||||
# NormalizeVideo(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
|
||||
transforms.Normalize(
|
||||
mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
target_video_len = 32
|
||||
frame_interval = 1
|
||||
@@ -582,7 +603,9 @@ if __name__ == "__main__":
|
||||
# print(start_frame_ind)
|
||||
# print(end_frame_ind)
|
||||
assert end_frame_ind - start_frame_ind >= target_video_len
|
||||
frame_indice = np.linspace(start_frame_ind, end_frame_ind - 1, target_video_len, dtype=int)
|
||||
frame_indice = np.linspace(
|
||||
start_frame_ind, end_frame_ind - 1, target_video_len, dtype=int
|
||||
)
|
||||
print(frame_indice)
|
||||
|
||||
select_vframes = vframes[frame_indice]
|
||||
@@ -593,7 +616,9 @@ if __name__ == "__main__":
|
||||
print(select_vframes_trans.shape)
|
||||
print(select_vframes_trans.dtype)
|
||||
|
||||
select_vframes_trans_int = ((select_vframes_trans * 0.5 + 0.5) * 255).to(dtype=torch.uint8)
|
||||
select_vframes_trans_int = ((select_vframes_trans * 0.5 + 0.5) * 255).to(
|
||||
dtype=torch.uint8
|
||||
)
|
||||
print(select_vframes_trans_int.dtype)
|
||||
print(select_vframes_trans_int.permute(0, 2, 3, 1).shape)
|
||||
|
||||
|
||||
+234
-112
@@ -1,41 +1,53 @@
|
||||
# !/bin/python3
|
||||
# isort: skip_file
|
||||
import argparse
|
||||
import math
|
||||
import os
|
||||
from fastvideo.utils.parallel_states import (
|
||||
initialize_sequence_parallel_state,
|
||||
destroy_sequence_parallel_group,
|
||||
get_sequence_parallel_state,
|
||||
nccl_info,
|
||||
)
|
||||
from fastvideo.utils.communications import sp_parallel_dataloader_wrapper, broadcast
|
||||
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.utils.validation import log_validation
|
||||
import time
|
||||
from collections import deque
|
||||
from copy import deepcopy
|
||||
|
||||
from torch.utils.data import DataLoader
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed.fsdp import ShardingStrategy
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
StateDictType,
|
||||
FullStateDictConfig,
|
||||
)
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
import json
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from fastvideo.utils.dataset_utils import LengthGroupedSampler
|
||||
import wandb
|
||||
from accelerate.utils import set_seed
|
||||
from tqdm.auto import tqdm
|
||||
from fastvideo.utils.fsdp_util import get_dit_fsdp_kwargs, apply_fsdp_checkpointing
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from fastvideo.utils.load import load_transformer
|
||||
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
|
||||
from copy import deepcopy
|
||||
from diffusers.optimization import get_scheduler
|
||||
from diffusers.utils import check_min_version
|
||||
from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_function
|
||||
import torch.distributed as dist
|
||||
from safetensors.torch import save_file
|
||||
from peft import LoraConfig
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.distributed.fsdp import ShardingStrategy
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
|
||||
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
|
||||
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
from fastvideo.utils.checkpoint import (resume_lora_optimizer, save_checkpoint, save_lora_checkpoint)
|
||||
from fastvideo.utils.communications import (broadcast, sp_parallel_dataloader_wrapper)
|
||||
from fastvideo.utils.dataset_utils import LengthGroupedSampler
|
||||
from fastvideo.utils.fsdp_util import (apply_fsdp_checkpointing, get_dit_fsdp_kwargs)
|
||||
from fastvideo.utils.load import load_transformer
|
||||
from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group, get_sequence_parallel_state,
|
||||
initialize_sequence_parallel_state)
|
||||
from fastvideo.utils.validation import log_validation
|
||||
from fastvideo.utils.checkpoint import (
|
||||
save_checkpoint,
|
||||
save_lora_checkpoint,
|
||||
resume_lora_optimizer,
|
||||
)
|
||||
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
check_min_version("0.31.0")
|
||||
import time
|
||||
from collections import deque
|
||||
|
||||
|
||||
def main_print(content):
|
||||
@@ -43,6 +55,31 @@ def main_print(content):
|
||||
print(content)
|
||||
|
||||
|
||||
def save_checkpoint(transformer, rank, output_dir, step):
|
||||
main_print(f"--> saving checkpoint at step {step}")
|
||||
with FSDP.state_dict_type(
|
||||
transformer,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
cpu_state = transformer.state_dict()
|
||||
# todo move to get_state_dict
|
||||
if rank <= 0:
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
config_dict = dict(transformer.config)
|
||||
if "dtype" in config_dict:
|
||||
del config_dict["dtype"] # TODO
|
||||
config_path = os.path.join(save_dir, "config.json")
|
||||
# save dict as json
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
main_print(f"--> checkpoint saved at step {step}")
|
||||
|
||||
|
||||
def reshard_fsdp(model):
|
||||
for m in FSDP.fsdp_modules(model):
|
||||
if m._has_params and m.sharding_strategy is not ShardingStrategy.NO_SHARD:
|
||||
@@ -51,15 +88,17 @@ def reshard_fsdp(model):
|
||||
|
||||
def get_norm(model_pred, norms, gradient_accumulation_steps):
|
||||
fro_norm = (
|
||||
torch.linalg.matrix_norm(model_pred, ord="fro") / # codespell:ignore
|
||||
gradient_accumulation_steps)
|
||||
largest_singular_value = (torch.linalg.matrix_norm(model_pred, ord=2) / gradient_accumulation_steps)
|
||||
torch.linalg.matrix_norm(model_pred, ord="fro") / gradient_accumulation_steps
|
||||
)
|
||||
largest_singular_value = (
|
||||
torch.linalg.matrix_norm(model_pred, ord=2) / gradient_accumulation_steps
|
||||
)
|
||||
absolute_mean = torch.mean(torch.abs(model_pred)) / gradient_accumulation_steps
|
||||
absolute_max = torch.max(torch.abs(model_pred)) / gradient_accumulation_steps
|
||||
dist.all_reduce(fro_norm, op=dist.ReduceOp.AVG)
|
||||
dist.all_reduce(largest_singular_value, op=dist.ReduceOp.AVG)
|
||||
dist.all_reduce(absolute_mean, op=dist.ReduceOp.AVG)
|
||||
norms["fro"] += torch.mean(fro_norm).item() # codespell:ignore
|
||||
norms["fro"] += torch.mean(fro_norm).item()
|
||||
norms["largest singular value"] += torch.mean(largest_singular_value).item()
|
||||
norms["absolute mean"] += absolute_mean.item()
|
||||
norms["absolute max"] += absolute_max.item()
|
||||
@@ -93,7 +132,7 @@ def distill_one_step(
|
||||
total_loss = 0.0
|
||||
optimizer.zero_grad()
|
||||
model_pred_norm = {
|
||||
"fro": 0.0, # codespell:ignore
|
||||
"fro": 0.0,
|
||||
"largest singular value": 0.0,
|
||||
"absolute mean": 0.0,
|
||||
"absolute max": 0.0,
|
||||
@@ -108,7 +147,9 @@ def distill_one_step(
|
||||
model_input = normalize_dit_input(model_type, latents)
|
||||
noise = torch.randn_like(model_input)
|
||||
bsz = model_input.shape[0]
|
||||
index = torch.randint(0, num_euler_timesteps, (bsz, ), device=model_input.device).long()
|
||||
index = torch.randint(
|
||||
0, num_euler_timesteps, (bsz,), device=model_input.device
|
||||
).long()
|
||||
if sp_size > 1:
|
||||
broadcast(index)
|
||||
# Add noise according to flow matching.
|
||||
@@ -119,7 +160,9 @@ def distill_one_step(
|
||||
timesteps = (sigmas * noise_scheduler.config.num_train_timesteps).view(-1)
|
||||
# if squeeze to [], unsqueeze to [1]
|
||||
|
||||
timesteps_prev = (sigmas_prev * noise_scheduler.config.num_train_timesteps).view(-1)
|
||||
timesteps_prev = (
|
||||
sigmas_prev * noise_scheduler.config.num_train_timesteps
|
||||
).view(-1)
|
||||
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
|
||||
# Predict the noise residual
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
@@ -131,13 +174,15 @@ def distill_one_step(
|
||||
"return_dict": False,
|
||||
}
|
||||
if hunyuan_teacher_disable_cfg:
|
||||
teacher_kwargs["guidance"] = torch.tensor([1000.0],
|
||||
device=noisy_model_input.device,
|
||||
dtype=torch.bfloat16)
|
||||
teacher_kwargs["guidance"] = torch.tensor(
|
||||
[1000.0], device=noisy_model_input.device, dtype=torch.bfloat16
|
||||
)
|
||||
model_pred = transformer(**teacher_kwargs)[0]
|
||||
|
||||
# if accelerator.is_main_process:
|
||||
model_pred, end_index = solver.euler_style_multiphase_pred(noisy_model_input, model_pred, index, multiphase)
|
||||
model_pred, end_index = solver.euler_style_multiphase_pred(
|
||||
noisy_model_input, model_pred, index, multiphase
|
||||
)
|
||||
with torch.no_grad():
|
||||
w = distill_cfg
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
@@ -160,7 +205,9 @@ def distill_one_step(
|
||||
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
teacher_output = uncond_teacher_output + w * (cond_teacher_output - uncond_teacher_output)
|
||||
teacher_output = cond_teacher_output + w * (
|
||||
cond_teacher_output - uncond_teacher_output
|
||||
)
|
||||
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
|
||||
|
||||
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
|
||||
@@ -183,26 +230,42 @@ def distill_one_step(
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
target, end_index = solver.euler_style_multiphase_pred(x_prev, target_pred, index, multiphase, True)
|
||||
target, end_index = solver.euler_style_multiphase_pred(
|
||||
x_prev, target_pred, index, multiphase, True
|
||||
)
|
||||
|
||||
huber_c = 0.001
|
||||
# loss = loss.mean()
|
||||
loss = (torch.mean(torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) - huber_c) /
|
||||
gradient_accumulation_steps)
|
||||
loss = (
|
||||
torch.mean(
|
||||
torch.sqrt((model_pred.float() - target.float()) ** 2 + huber_c ** 2)
|
||||
- huber_c
|
||||
)
|
||||
/ gradient_accumulation_steps
|
||||
)
|
||||
if pred_decay_weight > 0:
|
||||
if pred_decay_type == "l1":
|
||||
pred_decay_loss = (torch.mean(torch.sqrt(model_pred.float()**2)) * pred_decay_weight /
|
||||
gradient_accumulation_steps)
|
||||
pred_decay_loss = (
|
||||
torch.mean(torch.sqrt(model_pred.float() ** 2))
|
||||
* pred_decay_weight
|
||||
/ gradient_accumulation_steps
|
||||
)
|
||||
loss += pred_decay_loss
|
||||
elif pred_decay_type == "l2":
|
||||
# essnetially k2?
|
||||
pred_decay_loss = (torch.mean(model_pred.float()**2) * pred_decay_weight / gradient_accumulation_steps)
|
||||
pred_decay_loss = (
|
||||
torch.mean(model_pred.float() ** 2)
|
||||
* pred_decay_weight
|
||||
/ gradient_accumulation_steps
|
||||
)
|
||||
loss += pred_decay_loss
|
||||
else:
|
||||
assert NotImplementedError("pred_decay_type is not implemented")
|
||||
|
||||
# calculate model_pred norm and mean
|
||||
get_norm(model_pred.detach().float(), model_pred_norm, gradient_accumulation_steps)
|
||||
get_norm(
|
||||
model_pred.detach().float(), model_pred_norm, gradient_accumulation_steps
|
||||
)
|
||||
loss.backward()
|
||||
|
||||
avg_loss = loss.detach().clone()
|
||||
@@ -212,9 +275,13 @@ def distill_one_step(
|
||||
# update ema
|
||||
if ema_transformer is not None:
|
||||
reshard_fsdp(ema_transformer)
|
||||
for p_averaged, p_model in zip(ema_transformer.parameters(), transformer.parameters()):
|
||||
for p_averaged, p_model in zip(
|
||||
ema_transformer.parameters(), transformer.parameters()
|
||||
):
|
||||
with torch.no_grad():
|
||||
p_averaged.copy_(torch.lerp(p_averaged.detach(), p_model.detach(), 1 - ema_decay))
|
||||
p_averaged.copy_(
|
||||
torch.lerp(p_averaged.detach(), p_model.detach(), 1 - ema_decay)
|
||||
)
|
||||
|
||||
grad_norm = transformer.clip_grad_norm_(max_grad_norm)
|
||||
optimizer.step()
|
||||
@@ -245,7 +312,7 @@ def main(args):
|
||||
if rank <= 0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weights to half-precision
|
||||
# For mixed precision training we cast all non-trainable weigths to half-precision
|
||||
# as these weights are only used for inference, keeping weights in full precision is not required.
|
||||
|
||||
# Create model:
|
||||
@@ -277,8 +344,11 @@ def main(args):
|
||||
transformer.add_adapter(transformer_lora_config)
|
||||
|
||||
main_print(
|
||||
f" Total training parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M")
|
||||
main_print(f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}")
|
||||
f" Total training parameters = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e6} M"
|
||||
)
|
||||
main_print(
|
||||
f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}"
|
||||
)
|
||||
fsdp_kwargs, no_split_modules = get_dit_fsdp_kwargs(
|
||||
transformer,
|
||||
args.fsdp_sharding_startegy,
|
||||
@@ -294,26 +364,23 @@ def main(args):
|
||||
transformer._no_split_modules = no_split_modules
|
||||
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
|
||||
|
||||
transformer = FSDP(
|
||||
transformer,
|
||||
**fsdp_kwargs,
|
||||
)
|
||||
teacher_transformer = FSDP(
|
||||
teacher_transformer,
|
||||
**fsdp_kwargs,
|
||||
)
|
||||
transformer = FSDP(transformer, **fsdp_kwargs,)
|
||||
teacher_transformer = FSDP(teacher_transformer, **fsdp_kwargs,)
|
||||
if args.use_ema:
|
||||
ema_transformer = FSDP(
|
||||
ema_transformer,
|
||||
**fsdp_kwargs,
|
||||
)
|
||||
main_print("--> model loaded")
|
||||
ema_transformer = FSDP(ema_transformer, **fsdp_kwargs,)
|
||||
main_print(f"--> model loaded")
|
||||
|
||||
if args.gradient_checkpointing:
|
||||
apply_fsdp_checkpointing(transformer, no_split_modules, args.selective_checkpointing)
|
||||
apply_fsdp_checkpointing(teacher_transformer, no_split_modules, args.selective_checkpointing)
|
||||
apply_fsdp_checkpointing(
|
||||
transformer, no_split_modules, args.selective_checkpointing
|
||||
)
|
||||
apply_fsdp_checkpointing(
|
||||
teacher_transformer, no_split_modules, args.selective_checkpointing
|
||||
)
|
||||
if args.use_ema:
|
||||
apply_fsdp_checkpointing(ema_transformer, no_split_modules, args.selective_checkpointing)
|
||||
apply_fsdp_checkpointing(
|
||||
ema_transformer, no_split_modules, args.selective_checkpointing
|
||||
)
|
||||
# Set model as trainable.
|
||||
transformer.train()
|
||||
teacher_transformer.requires_grad_(False)
|
||||
@@ -321,7 +388,9 @@ def main(args):
|
||||
ema_transformer.requires_grad_(False)
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
|
||||
if args.scheduler_type == "pcm_linear_quadratic":
|
||||
linear_steps = int(noise_scheduler.config.num_train_timesteps * args.linear_range)
|
||||
linear_steps = int(
|
||||
noise_scheduler.config.num_train_timesteps * args.linear_range
|
||||
)
|
||||
sigmas = linear_quadratic_schedule(
|
||||
noise_scheduler.config.num_train_timesteps,
|
||||
args.linear_quadratic_threshold,
|
||||
@@ -349,8 +418,9 @@ def main(args):
|
||||
|
||||
init_steps = 0
|
||||
if args.resume_from_lora_checkpoint:
|
||||
transformer, optimizer, init_steps = resume_lora_optimizer(transformer, args.resume_from_lora_checkpoint,
|
||||
optimizer)
|
||||
transformer, optimizer, init_steps = resume_lora_optimizer(
|
||||
transformer, args.resume_from_lora_checkpoint, optimizer
|
||||
)
|
||||
main_print(f"optimizer: {optimizer}")
|
||||
|
||||
# todo add lr scheduler
|
||||
@@ -367,15 +437,20 @@ def main(args):
|
||||
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t, args.cfg)
|
||||
uncond_prompt_embed = train_dataset.uncond_prompt_embed
|
||||
uncond_prompt_mask = train_dataset.uncond_prompt_mask
|
||||
sampler = (LengthGroupedSampler(
|
||||
args.train_batch_size,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
lengths=train_dataset.lengths,
|
||||
group_frame=args.group_frame,
|
||||
group_resolution=args.group_resolution,
|
||||
) if (args.group_frame or args.group_resolution) else DistributedSampler(
|
||||
train_dataset, rank=rank, num_replicas=world_size, shuffle=False))
|
||||
sampler = (
|
||||
LengthGroupedSampler(
|
||||
args.train_batch_size,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
lengths=train_dataset.lengths,
|
||||
group_frame=args.group_frame,
|
||||
group_resolution=args.group_resolution,
|
||||
)
|
||||
if (args.group_frame or args.group_resolution)
|
||||
else DistributedSampler(
|
||||
train_dataset, rank=rank, num_replicas=world_size, shuffle=False
|
||||
)
|
||||
)
|
||||
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
@@ -388,7 +463,11 @@ def main(args):
|
||||
)
|
||||
|
||||
num_update_steps_per_epoch = math.ceil(
|
||||
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
|
||||
len(train_dataloader)
|
||||
/ args.gradient_accumulation_steps
|
||||
* args.sp_size
|
||||
/ args.train_sp_batch_size
|
||||
)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
if rank <= 0:
|
||||
@@ -396,14 +475,22 @@ def main(args):
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
# Train!
|
||||
total_batch_size = (world_size * args.gradient_accumulation_steps / args.sp_size * args.train_sp_batch_size)
|
||||
total_batch_size = (
|
||||
args.train_batch_size
|
||||
* world_size
|
||||
* args.gradient_accumulation_steps
|
||||
/ args.sp_size
|
||||
* args.train_sp_batch_size
|
||||
)
|
||||
main_print("***** Running training *****")
|
||||
main_print(f" Num examples = {len(train_dataset)}")
|
||||
main_print(f" Dataloader size = {len(train_dataloader)}")
|
||||
main_print(f" Num Epochs = {args.num_train_epochs}")
|
||||
main_print(f" Resume training from step {init_steps}")
|
||||
main_print(f" Instantaneous batch size per device = {args.train_batch_size}")
|
||||
main_print(f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}")
|
||||
main_print(
|
||||
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
|
||||
)
|
||||
main_print(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
|
||||
main_print(f" Total optimization steps = {args.max_train_steps}")
|
||||
main_print(
|
||||
@@ -486,12 +573,14 @@ def main(args):
|
||||
step_times.append(step_time)
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"loss": f"{loss:.4f}",
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
"grad_norm": grad_norm,
|
||||
"phases": num_phases,
|
||||
})
|
||||
progress_bar.set_postfix(
|
||||
{
|
||||
"loss": f"{loss:.4f}",
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
"grad_norm": grad_norm,
|
||||
"phases": num_phases,
|
||||
}
|
||||
)
|
||||
progress_bar.update(1)
|
||||
if rank <= 0:
|
||||
wandb.log(
|
||||
@@ -501,7 +590,7 @@ def main(args):
|
||||
"step_time": step_time,
|
||||
"avg_step_time": avg_step_time,
|
||||
"grad_norm": grad_norm,
|
||||
"pred_fro_norm": pred_norm["fro"], # codespell:ignore
|
||||
"pred_fro_norm": pred_norm["fro"],
|
||||
"pred_largest_singular_value": pred_norm["largest singular value"],
|
||||
"pred_absolute_mean": pred_norm["absolute mean"],
|
||||
"pred_absolute_max": pred_norm["absolute max"],
|
||||
@@ -511,7 +600,9 @@ def main(args):
|
||||
if step % args.checkpointing_steps == 0:
|
||||
if args.use_lora:
|
||||
# Save LoRA weights
|
||||
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, step)
|
||||
save_lora_checkpoint(
|
||||
transformer, optimizer, rank, args.output_dir, step
|
||||
)
|
||||
else:
|
||||
# Your existing checkpoint saving code
|
||||
if args.use_ema:
|
||||
@@ -549,7 +640,9 @@ def main(args):
|
||||
)
|
||||
|
||||
if args.use_lora:
|
||||
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps)
|
||||
save_lora_checkpoint(
|
||||
transformer, optimizer, rank, args.output_dir, args.max_train_steps
|
||||
)
|
||||
else:
|
||||
save_checkpoint(transformer, rank, args.output_dir, args.max_train_steps)
|
||||
|
||||
@@ -560,12 +653,12 @@ def main(args):
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument("--model_type", type=str, default="mochi", help="The type of model to train.")
|
||||
parser.add_argument(
|
||||
"--model_type", type=str, default="mochi", help="The type of model to train."
|
||||
)
|
||||
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--data_json_path", type=str, required=True)
|
||||
parser.add_argument("--num_height", type=int, default=480)
|
||||
parser.add_argument("--num_width", type=int, default=848)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument(
|
||||
"--dataloader_num_workers",
|
||||
@@ -579,7 +672,9 @@ if __name__ == "__main__":
|
||||
default=16,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
|
||||
parser.add_argument(
|
||||
"--num_latent_t", type=int, default=28, help="Number of latent timesteps."
|
||||
)
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
|
||||
@@ -601,7 +696,9 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--validation_steps", type=float, default=64)
|
||||
parser.add_argument("--log_validation", action="store_true")
|
||||
parser.add_argument("--tracker_project_name", type=str, default=None)
|
||||
parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
|
||||
parser.add_argument(
|
||||
"--seed", type=int, default=None, help="A seed for reproducible training."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
@@ -618,31 +715,39 @@ if __name__ == "__main__":
|
||||
"--checkpointing_steps",
|
||||
type=int,
|
||||
default=500,
|
||||
help=("Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
|
||||
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
|
||||
" training using `--resume_from_checkpoint`."),
|
||||
help=(
|
||||
"Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
|
||||
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
|
||||
" training using `--resume_from_checkpoint`."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--shift", type=float, default=1.0)
|
||||
parser.add_argument(
|
||||
"--resume_from_checkpoint",
|
||||
type=str,
|
||||
default=None,
|
||||
help=("Whether training should be resumed from a previous checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
|
||||
help=(
|
||||
"Whether training should be resumed from a previous checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--resume_from_lora_checkpoint",
|
||||
type=str,
|
||||
default=None,
|
||||
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
|
||||
help=(
|
||||
"Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
help=(
|
||||
"[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
|
||||
),
|
||||
)
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
@@ -677,7 +782,9 @@ if __name__ == "__main__":
|
||||
default=10,
|
||||
help="Number of steps for the warmup in the lr scheduler.",
|
||||
)
|
||||
parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.")
|
||||
parser.add_argument(
|
||||
"--max_grad_norm", default=1.0, type=float, help="Max gradient norm."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--gradient_checkpointing",
|
||||
action="store_true",
|
||||
@@ -687,8 +794,10 @@ if __name__ == "__main__":
|
||||
parser.add_argument(
|
||||
"--allow_tf32",
|
||||
action="store_true",
|
||||
help=("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
|
||||
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"),
|
||||
help=(
|
||||
"Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
|
||||
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mixed_precision",
|
||||
@@ -698,7 +807,8 @@ if __name__ == "__main__":
|
||||
help=(
|
||||
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
|
||||
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
|
||||
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."),
|
||||
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use_cpu_offload",
|
||||
@@ -720,8 +830,12 @@ if __name__ == "__main__":
|
||||
default=False,
|
||||
help="Whether to use LoRA for finetuning.",
|
||||
)
|
||||
parser.add_argument("--lora_alpha", type=int, default=256, help="Alpha parameter for LoRA.")
|
||||
parser.add_argument("--lora_rank", type=int, default=128, help="LoRA rank parameter. ")
|
||||
parser.add_argument(
|
||||
"--lora_alpha", type=int, default=256, help="Alpha parameter for LoRA."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lora_rank", type=int, default=128, help="LoRA rank parameter. "
|
||||
)
|
||||
parser.add_argument("--fsdp_sharding_startegy", default="full")
|
||||
|
||||
# lr_scheduler
|
||||
@@ -729,8 +843,10 @@ if __name__ == "__main__":
|
||||
"--lr_scheduler",
|
||||
type=str,
|
||||
default="constant",
|
||||
help=('The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
|
||||
' "constant", "constant_with_warmup"]'),
|
||||
help=(
|
||||
'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
|
||||
' "constant", "constant_with_warmup"]'
|
||||
),
|
||||
)
|
||||
parser.add_argument("--num_euler_timesteps", type=int, default=100)
|
||||
parser.add_argument(
|
||||
@@ -750,9 +866,13 @@ if __name__ == "__main__":
|
||||
action="store_true",
|
||||
help="Whether to apply the cfg_solver.",
|
||||
)
|
||||
parser.add_argument("--distill_cfg", type=float, default=3.0, help="Distillation coefficient.")
|
||||
parser.add_argument(
|
||||
"--distill_cfg", type=float, default=3.0, help="Distillation coefficient."
|
||||
)
|
||||
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
|
||||
parser.add_argument("--scheduler_type", type=str, default="pcm", help="The scheduler type to use.")
|
||||
parser.add_argument(
|
||||
"--scheduler_type", type=str, default="pcm", help="The scheduler type to use."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--linear_quadratic_threshold",
|
||||
type=float,
|
||||
@@ -765,7 +885,9 @@ if __name__ == "__main__":
|
||||
default=0.5,
|
||||
help="Range for linear quadratic scheduler.",
|
||||
)
|
||||
parser.add_argument("--weight_decay", type=float, default=0.001, help="Weight decay to apply.")
|
||||
parser.add_argument(
|
||||
"--weight_decay", type=float, default=0.001, help="Weight decay to apply."
|
||||
)
|
||||
parser.add_argument("--use_ema", action="store_true", help="Whether to use EMA.")
|
||||
parser.add_argument("--multi_phased_distill_schedule", type=str, default=None)
|
||||
parser.add_argument("--pred_decay_weight", type=float, default=0.0)
|
||||
|
||||
@@ -1,33 +1,63 @@
|
||||
from typing import Any, Dict, Optional, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from diffusers.utils import logging
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
|
||||
from diffusers.models.attention import JointTransformerBlock
|
||||
from diffusers.models.attention_processor import Attention, AttentionProcessor
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.models.normalization import AdaLayerNormContinuous
|
||||
from diffusers.utils import (
|
||||
USE_PEFT_BACKEND,
|
||||
is_torch_version,
|
||||
logging,
|
||||
scale_lora_layers,
|
||||
unscale_lora_layers,
|
||||
)
|
||||
from diffusers.models.embeddings import CombinedTimestepTextProjEmbeddings, PatchEmbed
|
||||
from diffusers.models.transformers.transformer_2d import Transformer2DModelOutput
|
||||
from diffusers.models.transformers.transformer_sd3 import SD3Transformer2DModel
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
class DiscriminatorHead(nn.Module):
|
||||
|
||||
def __init__(self, input_channel, output_channel=1):
|
||||
def __init__(self, input_channel, output_channel=1, args=None):
|
||||
super().__init__()
|
||||
inner_channel = 1024
|
||||
self.conv1 = nn.Sequential(
|
||||
nn.Conv2d(input_channel, inner_channel, 1, 1, 0),
|
||||
nn.GroupNorm(32, inner_channel),
|
||||
nn.LeakyReLU(inplace=True), # use LeakyReLu instead of GELU shown in the paper to save memory
|
||||
nn.LeakyReLU(
|
||||
inplace=True
|
||||
), # use LeakyReLu instead of GELU shown in the paper to save memory
|
||||
)
|
||||
self.conv2 = nn.Sequential(
|
||||
nn.Conv2d(inner_channel, inner_channel, 1, 1, 0),
|
||||
nn.GroupNorm(32, inner_channel),
|
||||
nn.LeakyReLU(inplace=True), # use LeakyReLu instead of GELU shown in the paper to save memory
|
||||
nn.LeakyReLU(
|
||||
inplace=True
|
||||
), # use LeakyReLu instead of GELU shown in the paper to save memory
|
||||
)
|
||||
|
||||
self.conv_out = nn.Conv2d(inner_channel, output_channel, 1, 1, 0)
|
||||
|
||||
vae_spatial_scale_factor = 8
|
||||
|
||||
self.patch_height = args.num_height // vae_spatial_scale_factor // 2
|
||||
self.patch_width = args.num_width // vae_spatial_scale_factor // 2
|
||||
print("## DiscriminatorHead: patch_height: ", self.patch_height)
|
||||
print("## DiscriminatorHead: patch_width: ", self.patch_width)
|
||||
|
||||
def forward(self, x):
|
||||
b, twh, c = x.shape
|
||||
t = twh // (30 * 53)
|
||||
x = x.view(-1, 30 * 53, c)
|
||||
|
||||
t = twh // (self.patch_height * self.patch_width)
|
||||
x = x.view(-1, self.patch_height * self.patch_width, c)
|
||||
x = x.permute(0, 2, 1)
|
||||
x = x.view(b * t, c, 30, 53)
|
||||
x = x.view(b * t, c, self.patch_height, self.patch_width)
|
||||
x = self.conv1(x)
|
||||
x = self.conv2(x) + x
|
||||
x = self.conv_out(x)
|
||||
@@ -35,29 +65,38 @@ class DiscriminatorHead(nn.Module):
|
||||
|
||||
|
||||
class Discriminator(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
stride=8,
|
||||
num_h_per_head=1,
|
||||
adapter_channel_dims=[3072],
|
||||
total_layers=48,
|
||||
total_layers = 48,
|
||||
args=None,
|
||||
):
|
||||
super().__init__()
|
||||
adapter_channel_dims = adapter_channel_dims * (total_layers // stride)
|
||||
self.stride = stride
|
||||
self.num_h_per_head = num_h_per_head
|
||||
self.head_num = len(adapter_channel_dims)
|
||||
self.heads = nn.ModuleList([
|
||||
nn.ModuleList([DiscriminatorHead(adapter_channel) for _ in range(self.num_h_per_head)])
|
||||
for adapter_channel in adapter_channel_dims
|
||||
])
|
||||
self.heads = nn.ModuleList(
|
||||
[
|
||||
nn.ModuleList(
|
||||
[
|
||||
DiscriminatorHead(
|
||||
adapter_channel,
|
||||
args=args
|
||||
)
|
||||
for _ in range(self.num_h_per_head)
|
||||
]
|
||||
)
|
||||
for adapter_channel in adapter_channel_dims
|
||||
]
|
||||
)
|
||||
|
||||
def forward(self, features):
|
||||
outputs = []
|
||||
|
||||
def create_custom_forward(module):
|
||||
|
||||
def custom_forward(*inputs):
|
||||
return module(*inputs)
|
||||
|
||||
@@ -74,3 +113,25 @@ class Discriminator(nn.Module):
|
||||
out = h(features[i])
|
||||
outputs.append(out)
|
||||
return outputs
|
||||
|
||||
|
||||
class DMDiscriminator(nn.Module):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
self.cls_pred_branch = nn.Sequential(
|
||||
nn.Conv2d(kernel_size=4, in_channels=1280, out_channels=1280, stride=2, padding=1), # 8x8 -> 4x4
|
||||
nn.GroupNorm(num_groups=32, num_channels=1280),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(kernel_size=4, in_channels=1280, out_channels=1280, stride=4, padding=0), # 4x4 -> 1x1
|
||||
nn.GroupNorm(num_groups=32, num_channels=1280),
|
||||
nn.SiLU(),
|
||||
nn.Conv2d(kernel_size=1, in_channels=1280, out_channels=1, stride=1, padding=0), # 1x1 -> 1x1
|
||||
)
|
||||
|
||||
self.cls_pred_branch.requires_grad_(True)
|
||||
|
||||
def forward(self, features):
|
||||
print("## features shape: ", features.shape)
|
||||
return self.cls_pred_branch(features)
|
||||
|
||||
+58
-32
@@ -3,10 +3,11 @@ from typing import Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
@@ -20,7 +21,7 @@ class PCMFMSchedulerOutput(BaseOutput):
|
||||
def extract_into_tensor(a, t, x_shape):
|
||||
b, *_ = t.shape
|
||||
out = a.gather(-1, t)
|
||||
return out.reshape(b, *((1, ) * (len(x_shape) - 1)))
|
||||
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
|
||||
|
||||
|
||||
class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
@@ -39,15 +40,20 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
):
|
||||
if linear_quadratic:
|
||||
linear_steps = int(num_train_timesteps * linear_range)
|
||||
sigmas = linear_quadratic_schedule(num_train_timesteps, linear_quadratic_threshold, linear_steps)
|
||||
sigmas = linear_quadratic_schedule(
|
||||
num_train_timesteps, linear_quadratic_threshold, linear_steps
|
||||
)
|
||||
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
|
||||
else:
|
||||
timesteps = np.linspace(1, num_train_timesteps, num_train_timesteps, dtype=np.float32)[::-1].copy()
|
||||
timesteps = np.linspace(
|
||||
1, num_train_timesteps, num_train_timesteps, dtype=np.float32
|
||||
)[::-1].copy()
|
||||
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
|
||||
sigmas = timesteps / num_train_timesteps
|
||||
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
|
||||
self.euler_timesteps = (np.arange(1, pcm_timesteps + 1) *
|
||||
(num_train_timesteps // pcm_timesteps)).round().astype(np.int64) - 1
|
||||
self.euler_timesteps = (
|
||||
np.arange(1, pcm_timesteps + 1) * (num_train_timesteps // pcm_timesteps)
|
||||
).round().astype(np.int64) - 1
|
||||
self.sigmas = sigmas.numpy()[::-1][self.euler_timesteps]
|
||||
self.sigmas = torch.from_numpy((self.sigmas[::-1].copy()))
|
||||
self.timesteps = self.sigmas * num_train_timesteps
|
||||
@@ -112,7 +118,9 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
def _sigma_to_t(self, sigma):
|
||||
return sigma * self.config.num_train_timesteps
|
||||
|
||||
def set_timesteps(self, num_inference_steps: int, device: Union[str, torch.device] = None):
|
||||
def set_timesteps(
|
||||
self, num_inference_steps: int, device: Union[str, torch.device] = None
|
||||
):
|
||||
"""
|
||||
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
|
||||
@@ -123,14 +131,18 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
"""
|
||||
self.num_inference_steps = num_inference_steps
|
||||
inference_indices = np.linspace(0, self.config.pcm_timesteps, num=num_inference_steps, endpoint=False)
|
||||
inference_indices = np.linspace(
|
||||
0, self.config.pcm_timesteps, num=num_inference_steps, endpoint=False
|
||||
)
|
||||
inference_indices = np.floor(inference_indices).astype(np.int64)
|
||||
inference_indices = torch.from_numpy(inference_indices).long()
|
||||
|
||||
self.sigmas_ = self.sigmas[inference_indices]
|
||||
timesteps = self.sigmas_ * self.config.num_train_timesteps
|
||||
self.timesteps = timesteps.to(device=device)
|
||||
self.sigmas_ = torch.cat([self.sigmas_, torch.zeros(1, device=self.sigmas_.device)])
|
||||
self.sigmas_ = torch.cat(
|
||||
[self.sigmas_, torch.zeros(1, device=self.sigmas_.device)]
|
||||
)
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
@@ -192,11 +204,18 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
returned, otherwise a tuple is returned where the first element is the sample tensor.
|
||||
"""
|
||||
|
||||
if (isinstance(timestep, int) or isinstance(timestep, torch.IntTensor)
|
||||
or isinstance(timestep, torch.LongTensor)):
|
||||
raise ValueError(("Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||
" one of the `scheduler.timesteps` as a timestep."), )
|
||||
if (
|
||||
isinstance(timestep, int)
|
||||
or isinstance(timestep, torch.IntTensor)
|
||||
or isinstance(timestep, torch.LongTensor)
|
||||
):
|
||||
raise ValueError(
|
||||
(
|
||||
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||
" one of the `scheduler.timesteps` as a timestep."
|
||||
),
|
||||
)
|
||||
|
||||
if self.step_index is None:
|
||||
self._init_step_index(timestep)
|
||||
@@ -214,7 +233,7 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
self._step_index += 1
|
||||
|
||||
if not return_dict:
|
||||
return (prev_sample, )
|
||||
return (prev_sample,)
|
||||
|
||||
return PCMFMSchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
@@ -223,14 +242,16 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
|
||||
|
||||
class EulerSolver:
|
||||
|
||||
def __init__(self, sigmas, timesteps=1000, euler_timesteps=50):
|
||||
self.step_ratio = timesteps // euler_timesteps
|
||||
self.euler_timesteps = (np.arange(1, euler_timesteps + 1) * self.step_ratio).round().astype(np.int64) - 1
|
||||
self.euler_timesteps = (
|
||||
np.arange(1, euler_timesteps + 1) * self.step_ratio
|
||||
).round().astype(np.int64) - 1
|
||||
self.euler_timesteps_prev = np.asarray([0] + self.euler_timesteps[:-1].tolist())
|
||||
self.sigmas = sigmas[self.euler_timesteps]
|
||||
self.sigmas_prev = np.asarray([sigmas[0]] +
|
||||
sigmas[self.euler_timesteps[:-1]].tolist()) # either use sigma0 or 0
|
||||
self.sigmas_prev = np.asarray(
|
||||
[sigmas[0]] + sigmas[self.euler_timesteps[:-1]].tolist()
|
||||
) # either use sigma0 or 0
|
||||
|
||||
self.euler_timesteps = torch.from_numpy(self.euler_timesteps).long()
|
||||
self.euler_timesteps_prev = torch.from_numpy(self.euler_timesteps_prev).long()
|
||||
@@ -247,22 +268,25 @@ class EulerSolver:
|
||||
|
||||
def euler_step(self, sample, model_pred, timestep_index):
|
||||
sigma = extract_into_tensor(self.sigmas, timestep_index, model_pred.shape)
|
||||
sigma_prev = extract_into_tensor(self.sigmas_prev, timestep_index, model_pred.shape)
|
||||
sigma_prev = extract_into_tensor(
|
||||
self.sigmas_prev, timestep_index, model_pred.shape
|
||||
)
|
||||
x_prev = sample + (sigma_prev - sigma) * model_pred
|
||||
return x_prev
|
||||
|
||||
def euler_style_multiphase_pred(
|
||||
self,
|
||||
sample,
|
||||
model_pred,
|
||||
timestep_index,
|
||||
multiphase,
|
||||
is_target=False,
|
||||
self, sample, model_pred, timestep_index, multiphase, is_target=False,
|
||||
):
|
||||
inference_indices = np.linspace(0, len(self.euler_timesteps), num=multiphase, endpoint=False)
|
||||
inference_indices = np.linspace(
|
||||
0, len(self.euler_timesteps), num=multiphase, endpoint=False
|
||||
)
|
||||
inference_indices = np.floor(inference_indices).astype(np.int64)
|
||||
inference_indices = (torch.from_numpy(inference_indices).long().to(self.euler_timesteps.device))
|
||||
expanded_timestep_index = timestep_index.unsqueeze(1).expand(-1, inference_indices.size(0))
|
||||
inference_indices = (
|
||||
torch.from_numpy(inference_indices).long().to(self.euler_timesteps.device)
|
||||
)
|
||||
expanded_timestep_index = timestep_index.unsqueeze(1).expand(
|
||||
-1, inference_indices.size(0)
|
||||
)
|
||||
valid_indices_mask = expanded_timestep_index >= inference_indices
|
||||
last_valid_index = valid_indices_mask.flip(dims=[1]).long().argmax(dim=1)
|
||||
last_valid_index = inference_indices.size(0) - 1 - last_valid_index
|
||||
@@ -272,7 +296,9 @@ class EulerSolver:
|
||||
sigma = extract_into_tensor(self.sigmas_prev, timestep_index, sample.shape)
|
||||
else:
|
||||
sigma = extract_into_tensor(self.sigmas, timestep_index, sample.shape)
|
||||
sigma_prev = extract_into_tensor(self.sigmas_prev, timestep_index_end, sample.shape)
|
||||
sigma_prev = extract_into_tensor(
|
||||
self.sigmas_prev, timestep_index_end, sample.shape
|
||||
)
|
||||
x_prev = sample + (sigma_prev - sigma) * model_pred
|
||||
|
||||
return x_prev, timestep_index_end
|
||||
|
||||
+247
-112
@@ -1,43 +1,98 @@
|
||||
# !/bin/python3
|
||||
# isort: skip_file
|
||||
import argparse
|
||||
from email.policy import strict
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
from fastvideo.utils.parallel_states import (
|
||||
initialize_sequence_parallel_state,
|
||||
destroy_sequence_parallel_group,
|
||||
get_sequence_parallel_state,
|
||||
nccl_info,
|
||||
)
|
||||
from fastvideo.utils.communications import sp_parallel_dataloader_wrapper, broadcast
|
||||
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.utils.validation import log_validation
|
||||
import time
|
||||
from collections import deque
|
||||
from copy import deepcopy
|
||||
|
||||
from torch.utils.data import DataLoader
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
StateDictType,
|
||||
FullStateDictConfig,
|
||||
)
|
||||
from fastvideo.utils.load import load_transformer
|
||||
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
import json
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from fastvideo.utils.dataset_utils import LengthGroupedSampler
|
||||
import wandb
|
||||
from accelerate.utils import set_seed
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.optimization import get_scheduler
|
||||
from diffusers.utils import check_min_version
|
||||
from peft import LoraConfig
|
||||
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
|
||||
from fastvideo.utils.fsdp_util import (
|
||||
get_dit_fsdp_kwargs,
|
||||
apply_fsdp_checkpointing,
|
||||
get_discriminator_fsdp_kwargs,
|
||||
)
|
||||
import diffusers
|
||||
from diffusers import (
|
||||
FlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
from fastvideo.distill.discriminator import Discriminator
|
||||
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
|
||||
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
from fastvideo.utils.checkpoint import (resume_lora_optimizer, resume_training_generator_discriminator, save_checkpoint,
|
||||
save_lora_checkpoint)
|
||||
from fastvideo.utils.communications import (broadcast, sp_parallel_dataloader_wrapper)
|
||||
from fastvideo.utils.dataset_utils import LengthGroupedSampler
|
||||
from fastvideo.utils.fsdp_util import (apply_fsdp_checkpointing, get_discriminator_fsdp_kwargs, get_dit_fsdp_kwargs)
|
||||
from fastvideo.utils.load import load_transformer
|
||||
from copy import deepcopy
|
||||
from diffusers.optimization import get_scheduler
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
|
||||
from diffusers.utils import check_min_version
|
||||
from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_function
|
||||
import torch.distributed as dist
|
||||
from peft import LoraConfig
|
||||
from torch.distributed.fsdp import (
|
||||
FullyShardedDataParallel as FSDP,
|
||||
)
|
||||
from fastvideo.utils.checkpoint import (
|
||||
save_lora_checkpoint,
|
||||
resume_lora_optimizer,
|
||||
resume_training,
|
||||
save_checkpoint_generator_discriminator,
|
||||
resume_training_generator_discriminator,
|
||||
)
|
||||
# from fastvideo.utils.checkpoint import save_checkpoint
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group, get_sequence_parallel_state,
|
||||
initialize_sequence_parallel_state)
|
||||
from fastvideo.utils.validation import log_validation
|
||||
from torch.distributed.fsdp import FullOptimStateDictConfig
|
||||
from safetensors.torch import save_file
|
||||
|
||||
def save_checkpoint(model, rank, output_dir, step, discriminator=False):
|
||||
with FSDP.state_dict_type(
|
||||
model,
|
||||
StateDictType.FULL_STATE_DICT,
|
||||
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True),
|
||||
):
|
||||
cpu_state = model.state_dict()
|
||||
|
||||
# todo move to get_state_dict
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
if rank <= 0 and not discriminator:
|
||||
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
config_dict = dict(model.config)
|
||||
config_path = os.path.join(save_dir, "config.json")
|
||||
# save dict as json
|
||||
with open(config_path, "w") as f:
|
||||
json.dump(config_dict, f, indent=4)
|
||||
else:
|
||||
weight_path = os.path.join(save_dir, "discriminator_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
check_min_version("0.31.0")
|
||||
import time
|
||||
from collections import deque
|
||||
|
||||
|
||||
def gan_d_loss(
|
||||
@@ -49,7 +104,7 @@ def gan_d_loss(
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
weight,
|
||||
discriminator_head_stride,
|
||||
discriminator_head_stride
|
||||
):
|
||||
loss = 0.0
|
||||
# collate sample_fake and sample_real
|
||||
@@ -76,8 +131,10 @@ def gan_d_loss(
|
||||
fake_outputs = discriminator(fake_features)
|
||||
real_outputs = discriminator(real_features)
|
||||
for fake_output, real_output in zip(fake_outputs, real_outputs):
|
||||
loss += (torch.mean(weight * torch.relu(fake_output.float() + 1)) + torch.mean(
|
||||
weight * torch.relu(1 - real_output.float()))) / (discriminator.head_num * discriminator.num_h_per_head)
|
||||
loss += (
|
||||
torch.mean(weight * torch.relu(fake_output.float() + 1))
|
||||
+ torch.mean(weight * torch.relu(1 - real_output.float()))
|
||||
) / (discriminator.head_num * discriminator.num_h_per_head)
|
||||
return loss
|
||||
|
||||
|
||||
@@ -89,7 +146,7 @@ def gan_g_loss(
|
||||
encoder_hidden_states,
|
||||
encoder_attention_mask,
|
||||
weight,
|
||||
discriminator_head_stride,
|
||||
discriminator_head_stride
|
||||
):
|
||||
loss = 0.0
|
||||
features = teacher_transformer(
|
||||
@@ -101,10 +158,13 @@ def gan_g_loss(
|
||||
output_features_stride=discriminator_head_stride,
|
||||
return_dict=False,
|
||||
)[1]
|
||||
fake_outputs = discriminator(features, )
|
||||
fake_outputs = discriminator(
|
||||
features,
|
||||
)
|
||||
for fake_output in fake_outputs:
|
||||
loss += torch.mean(
|
||||
weight * torch.relu(1 - fake_output.float())) / (discriminator.head_num * discriminator.num_h_per_head)
|
||||
loss += torch.mean(weight * torch.relu(1 - fake_output.float())) / (
|
||||
discriminator.head_num * discriminator.num_h_per_head
|
||||
)
|
||||
return loss
|
||||
|
||||
|
||||
@@ -129,7 +189,7 @@ def distill_one_step_adv(
|
||||
not_apply_cfg_solver,
|
||||
distill_cfg,
|
||||
adv_weight,
|
||||
discriminator_head_stride,
|
||||
discriminator_head_stride
|
||||
):
|
||||
optimizer.zero_grad()
|
||||
discriminator_optimizer.zero_grad()
|
||||
@@ -143,7 +203,9 @@ def distill_one_step_adv(
|
||||
model_input = normalize_dit_input(model_type, latents)
|
||||
noise = torch.randn_like(model_input)
|
||||
bsz = model_input.shape[0]
|
||||
index = torch.randint(0, num_euler_timesteps, (bsz, ), device=model_input.device).long()
|
||||
index = torch.randint(
|
||||
0, num_euler_timesteps, (bsz,), device=model_input.device
|
||||
).long()
|
||||
if sp_size > 1:
|
||||
broadcast(index)
|
||||
# Add noise according to flow matching.
|
||||
@@ -168,8 +230,11 @@ def distill_one_step_adv(
|
||||
)[0]
|
||||
|
||||
# if accelerator.is_main_process:
|
||||
model_pred, end_index = solver.euler_style_multiphase_pred(noisy_model_input, model_pred, index, multiphase)
|
||||
model_pred, end_index = solver.euler_style_multiphase_pred(
|
||||
noisy_model_input, model_pred, index, multiphase
|
||||
)
|
||||
|
||||
weighting = 1.0
|
||||
# # simplified flow matching aka 0-rectified flow matching loss
|
||||
# # target = model_input - noise
|
||||
# target = model_input
|
||||
@@ -178,13 +243,14 @@ def distill_one_step_adv(
|
||||
adv_index[i] = torch.randint(
|
||||
end_index[i].item(),
|
||||
end_index[i].item() + num_euler_timesteps // multiphase,
|
||||
(1, ),
|
||||
(1,),
|
||||
dtype=end_index.dtype,
|
||||
device=end_index.device,
|
||||
)
|
||||
|
||||
sigmas_end = extract_into_tensor(solver.sigmas_prev, end_index, model_input.shape)
|
||||
sigmas_adv = extract_into_tensor(solver.sigmas_prev, adv_index, model_input.shape)
|
||||
timesteps_end = (sigmas_end * noise_scheduler.config.num_train_timesteps).view(-1)
|
||||
timesteps_adv = (sigmas_adv * noise_scheduler.config.num_train_timesteps).view(-1)
|
||||
|
||||
with torch.no_grad():
|
||||
@@ -209,7 +275,9 @@ def distill_one_step_adv(
|
||||
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
teacher_output = cond_teacher_output + w * (cond_teacher_output - uncond_teacher_output)
|
||||
teacher_output = cond_teacher_output + w * (
|
||||
cond_teacher_output - uncond_teacher_output
|
||||
)
|
||||
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
|
||||
|
||||
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
|
||||
@@ -223,14 +291,22 @@ def distill_one_step_adv(
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
target, end_index = solver.euler_style_multiphase_pred(x_prev, target_pred, index, multiphase, True)
|
||||
target, end_index = solver.euler_style_multiphase_pred(
|
||||
x_prev, target_pred, index, multiphase, True
|
||||
)
|
||||
|
||||
real_adv = ((1 - sigmas_adv) * target + (sigmas_adv - sigmas_end) * torch.randn_like(target)) / (1 - sigmas_end)
|
||||
fake_adv = ((1 - sigmas_adv) * model_pred +
|
||||
(sigmas_adv - sigmas_end) * torch.randn_like(model_pred)) / (1 - sigmas_end)
|
||||
real_adv = (
|
||||
(1 - sigmas_adv) * target + (sigmas_adv - sigmas_end) * torch.randn_like(target)
|
||||
) / (1 - sigmas_end)
|
||||
fake_adv = (
|
||||
(1 - sigmas_adv) * model_pred
|
||||
+ (sigmas_adv - sigmas_end) * torch.randn_like(model_pred)
|
||||
) / (1 - sigmas_end)
|
||||
|
||||
huber_c = 0.001
|
||||
g_loss = torch.mean(torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) - huber_c)
|
||||
g_loss = torch.mean(
|
||||
torch.sqrt((model_pred.float() - target.float()) ** 2 + huber_c**2) - huber_c
|
||||
)
|
||||
discriminator.requires_grad_(False)
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
g_gan_loss = adv_weight * gan_g_loss(
|
||||
@@ -241,7 +317,7 @@ def distill_one_step_adv(
|
||||
encoder_hidden_states.float(),
|
||||
encoder_attention_mask,
|
||||
1.0,
|
||||
discriminator_head_stride,
|
||||
discriminator_head_stride
|
||||
)
|
||||
g_loss += g_gan_loss
|
||||
g_loss.backward()
|
||||
@@ -299,7 +375,7 @@ def main(args):
|
||||
if rank <= 0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weights to half-precision
|
||||
# For mixed precision training we cast all non-trainable weigths to half-precision
|
||||
# as these weights are only used for inference, keeping weights in full precision is not required.
|
||||
|
||||
# Create model:
|
||||
@@ -313,10 +389,7 @@ def main(args):
|
||||
torch.float32 if args.master_weight_type == "fp32" else torch.bfloat16,
|
||||
)
|
||||
teacher_transformer = deepcopy(transformer)
|
||||
discriminator = Discriminator(
|
||||
args.discriminator_head_stride,
|
||||
total_layers=48 if args.model_type == "mochi" else 40,
|
||||
)
|
||||
discriminator = Discriminator(args.discriminator_head_stride, total_layers = 48 if args.model_type =="mochi" else 40)
|
||||
|
||||
if args.use_lora:
|
||||
transformer.requires_grad_(False)
|
||||
@@ -335,7 +408,9 @@ def main(args):
|
||||
main_print(
|
||||
f" Total discriminator parameters = {sum(p.numel() for p in discriminator.parameters() if p.requires_grad) / 1e6} M"
|
||||
)
|
||||
main_print(f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}")
|
||||
main_print(
|
||||
f"--> Initializing FSDP with sharding strategy: {args.fsdp_sharding_startegy}"
|
||||
)
|
||||
fsdp_kwargs, no_split_modules = get_dit_fsdp_kwargs(
|
||||
transformer,
|
||||
args.fsdp_sharding_startegy,
|
||||
@@ -364,17 +439,23 @@ def main(args):
|
||||
discriminator,
|
||||
**discriminator_fsdp_kwargs,
|
||||
)
|
||||
main_print("--> model loaded")
|
||||
main_print(f"--> model loaded")
|
||||
|
||||
if args.gradient_checkpointing:
|
||||
apply_fsdp_checkpointing(transformer, no_split_modules, args.selective_checkpointing)
|
||||
apply_fsdp_checkpointing(teacher_transformer, no_split_modules, args.selective_checkpointing)
|
||||
apply_fsdp_checkpointing(
|
||||
transformer, no_split_modules, args.selective_checkpointing
|
||||
)
|
||||
apply_fsdp_checkpointing(
|
||||
teacher_transformer, no_split_modules, args.selective_checkpointing
|
||||
)
|
||||
# Set model as trainable.
|
||||
transformer.train()
|
||||
teacher_transformer.requires_grad_(False)
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
|
||||
if args.scheduler_type == "pcm_linear_quadratic":
|
||||
sigmas = linear_quadratic_schedule(noise_scheduler.config.num_train_timesteps, args.linear_quadratic_threshold)
|
||||
sigmas = linear_quadratic_schedule(
|
||||
noise_scheduler.config.num_train_timesteps, args.linear_quadratic_threshold
|
||||
)
|
||||
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
|
||||
else:
|
||||
sigmas = noise_scheduler.sigmas
|
||||
@@ -391,7 +472,7 @@ def main(args):
|
||||
params_to_optimize,
|
||||
lr=args.learning_rate,
|
||||
betas=(0.9, 0.999),
|
||||
weight_decay=args.weight_decay,
|
||||
weight_decay=1e-3,
|
||||
eps=1e-8,
|
||||
)
|
||||
|
||||
@@ -399,14 +480,15 @@ def main(args):
|
||||
discriminator.parameters(),
|
||||
lr=args.discriminator_learning_rate,
|
||||
betas=(0, 0.999),
|
||||
weight_decay=args.weight_decay,
|
||||
weight_decay=1e-3,
|
||||
eps=1e-8,
|
||||
)
|
||||
|
||||
init_steps = 0
|
||||
if args.resume_from_lora_checkpoint:
|
||||
transformer, optimizer, init_steps = resume_lora_optimizer(transformer, args.resume_from_lora_checkpoint,
|
||||
optimizer)
|
||||
transformer, optimizer, init_steps = resume_lora_optimizer(
|
||||
transformer, args.resume_from_lora_checkpoint, optimizer
|
||||
)
|
||||
elif args.resume_from_checkpoint:
|
||||
(
|
||||
transformer,
|
||||
@@ -438,15 +520,20 @@ def main(args):
|
||||
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t, args.cfg)
|
||||
uncond_prompt_embed = train_dataset.uncond_prompt_embed
|
||||
uncond_prompt_mask = train_dataset.uncond_prompt_mask
|
||||
sampler = (LengthGroupedSampler(
|
||||
args.train_batch_size,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
lengths=train_dataset.lengths,
|
||||
group_frame=args.group_frame,
|
||||
group_resolution=args.group_resolution,
|
||||
) if (args.group_frame or args.group_resolution) else DistributedSampler(
|
||||
train_dataset, rank=rank, num_replicas=world_size, shuffle=False))
|
||||
sampler = (
|
||||
LengthGroupedSampler(
|
||||
args.train_batch_size,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
lengths=train_dataset.lengths,
|
||||
group_frame=args.group_frame,
|
||||
group_resolution=args.group_resolution,
|
||||
)
|
||||
if (args.group_frame or args.group_resolution)
|
||||
else DistributedSampler(
|
||||
train_dataset, rank=rank, num_replicas=world_size, shuffle=False
|
||||
)
|
||||
)
|
||||
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
@@ -459,7 +546,11 @@ def main(args):
|
||||
)
|
||||
assert args.gradient_accumulation_steps == 1
|
||||
num_update_steps_per_epoch = math.ceil(
|
||||
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
|
||||
len(train_dataloader)
|
||||
/ args.gradient_accumulation_steps
|
||||
* args.sp_size
|
||||
/ args.train_sp_batch_size
|
||||
)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
if rank <= 0:
|
||||
@@ -467,14 +558,22 @@ def main(args):
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
# Train!
|
||||
total_batch_size = (world_size * args.gradient_accumulation_steps / args.sp_size * args.train_sp_batch_size)
|
||||
total_batch_size = (
|
||||
args.train_batch_size
|
||||
* world_size
|
||||
* args.gradient_accumulation_steps
|
||||
/ args.sp_size
|
||||
* args.train_sp_batch_size
|
||||
)
|
||||
main_print("***** Running training *****")
|
||||
main_print(f" Num examples = {len(train_dataset)}")
|
||||
main_print(f" Dataloader size = {len(train_dataloader)}")
|
||||
main_print(f" Num Epochs = {args.num_train_epochs}")
|
||||
main_print(f" Resume training from step {init_steps}")
|
||||
main_print(f" Instantaneous batch size per device = {args.train_batch_size}")
|
||||
main_print(f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}")
|
||||
main_print(
|
||||
f" Total train batch size (w. data & sequence parallel, accumulation) = {total_batch_size}"
|
||||
)
|
||||
main_print(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
|
||||
main_print(f" Total optimization steps = {args.max_train_steps}")
|
||||
main_print(
|
||||
@@ -499,7 +598,6 @@ def main(args):
|
||||
)
|
||||
|
||||
step_times = deque(maxlen=100)
|
||||
|
||||
# log_validation(args, transformer, device,
|
||||
# torch.bfloat16, 0, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,ema=False)
|
||||
def get_num_phases(multi_phased_distill_schedule, step):
|
||||
@@ -511,7 +609,6 @@ def main(args):
|
||||
if step <= int(phase_step):
|
||||
return int(phase)
|
||||
return phase
|
||||
|
||||
for i in range(init_steps):
|
||||
_ = next(loader)
|
||||
for step in range(init_steps + 1, args.max_train_steps + 1):
|
||||
@@ -544,20 +641,22 @@ def main(args):
|
||||
args.not_apply_cfg_solver,
|
||||
args.distill_cfg,
|
||||
args.adv_weight,
|
||||
args.discriminator_head_stride,
|
||||
args.discriminator_head_stride
|
||||
)
|
||||
|
||||
step_time = time.time() - start_time
|
||||
step_times.append(step_time)
|
||||
avg_step_time = sum(step_times) / len(step_times)
|
||||
|
||||
progress_bar.set_postfix({
|
||||
"g_loss": f"{generator_loss:.4f}",
|
||||
"d_loss": f"{discriminator_loss:.4f}",
|
||||
"g_grad_norm": generator_grad_norm,
|
||||
"d_grad_norm": discriminator_grad_norm,
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
})
|
||||
progress_bar.set_postfix(
|
||||
{
|
||||
"g_loss": f"{generator_loss:.4f}",
|
||||
"d_loss": f"{discriminator_loss:.4f}",
|
||||
"g_grad_norm": generator_grad_norm,
|
||||
"d_grad_norm": discriminator_grad_norm,
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
}
|
||||
)
|
||||
progress_bar.update(1)
|
||||
if rank <= 0:
|
||||
wandb.log(
|
||||
@@ -576,10 +675,12 @@ def main(args):
|
||||
main_print(f"--> saving checkpoint at step {step}")
|
||||
if args.use_lora:
|
||||
# Save LoRA weights
|
||||
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, step)
|
||||
save_lora_checkpoint(
|
||||
transformer, optimizer, rank, args.output_dir, step
|
||||
)
|
||||
else:
|
||||
# Your existing checkpoint saving code
|
||||
# TODO
|
||||
# TODO
|
||||
# save_checkpoint_generator_discriminator(
|
||||
# transformer,
|
||||
# optimizer,
|
||||
@@ -589,7 +690,7 @@ def main(args):
|
||||
# args.output_dir,
|
||||
# step,
|
||||
# )
|
||||
save_checkpoint(transformer, rank, args.output_dir, step)
|
||||
save_checkpoint(transformer, rank, args.output_dir, step, discriminator)
|
||||
main_print(f"--> checkpoint saved at step {step}")
|
||||
dist.barrier()
|
||||
if args.log_validation and step % args.validation_steps == 0:
|
||||
@@ -606,9 +707,12 @@ def main(args):
|
||||
linear_range=args.linear_range,
|
||||
ema=False,
|
||||
)
|
||||
|
||||
|
||||
if args.use_lora:
|
||||
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, args.max_train_steps)
|
||||
save_lora_checkpoint(
|
||||
transformer, optimizer, rank, args.output_dir, args.max_train_steps
|
||||
)
|
||||
else:
|
||||
save_checkpoint(transformer, rank, args.output_dir, args.max_train_steps)
|
||||
|
||||
@@ -619,7 +723,9 @@ def main(args):
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser()
|
||||
|
||||
parser.add_argument("--model_type", type=str, default="mochi", help="The type of model to train.")
|
||||
parser.add_argument(
|
||||
"--model_type", type=str, default="mochi", help="The type of model to train."
|
||||
)
|
||||
# dataset & dataloader
|
||||
parser.add_argument("--data_json_path", type=str, required=True)
|
||||
parser.add_argument("--num_height", type=int, default=480)
|
||||
@@ -637,7 +743,9 @@ if __name__ == "__main__":
|
||||
default=16,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
|
||||
parser.add_argument(
|
||||
"--num_latent_t", type=int, default=28, help="Number of latent timesteps."
|
||||
)
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
|
||||
@@ -656,7 +764,9 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--validation_steps", type=float, default=64)
|
||||
parser.add_argument("--log_validation", action="store_true")
|
||||
parser.add_argument("--tracker_project_name", type=str, default=None)
|
||||
parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.")
|
||||
parser.add_argument(
|
||||
"--seed", type=int, default=None, help="A seed for reproducible training."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
@@ -673,9 +783,11 @@ if __name__ == "__main__":
|
||||
"--checkpointing_steps",
|
||||
type=int,
|
||||
default=500,
|
||||
help=("Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
|
||||
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
|
||||
" training using `--resume_from_checkpoint`."),
|
||||
help=(
|
||||
"Save a checkpoint of the training state every X updates. These checkpoints can be used both as final"
|
||||
" checkpoints in case they are better than the last checkpoint, and are also suitable for resuming"
|
||||
" training using `--resume_from_checkpoint`."
|
||||
),
|
||||
)
|
||||
parser.add_argument("--validation_prompt_dir", type=str)
|
||||
parser.add_argument("--shift", type=float, default=1.0)
|
||||
@@ -683,22 +795,28 @@ if __name__ == "__main__":
|
||||
"--resume_from_checkpoint",
|
||||
type=str,
|
||||
default=None,
|
||||
help=("Whether training should be resumed from a previous checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
|
||||
help=(
|
||||
"Whether training should be resumed from a previous checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--resume_from_lora_checkpoint",
|
||||
type=str,
|
||||
default=None,
|
||||
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
|
||||
help=(
|
||||
"Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
help=(
|
||||
"[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
|
||||
),
|
||||
)
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
@@ -733,7 +851,9 @@ if __name__ == "__main__":
|
||||
default=10,
|
||||
help="Number of steps for the warmup in the lr scheduler.",
|
||||
)
|
||||
parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.")
|
||||
parser.add_argument(
|
||||
"--max_grad_norm", default=1.0, type=float, help="Max gradient norm."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--gradient_checkpointing",
|
||||
action="store_true",
|
||||
@@ -743,8 +863,10 @@ if __name__ == "__main__":
|
||||
parser.add_argument(
|
||||
"--allow_tf32",
|
||||
action="store_true",
|
||||
help=("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
|
||||
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"),
|
||||
help=(
|
||||
"Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
|
||||
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--mixed_precision",
|
||||
@@ -754,7 +876,8 @@ if __name__ == "__main__":
|
||||
help=(
|
||||
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
|
||||
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
|
||||
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."),
|
||||
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use_cpu_offload",
|
||||
@@ -776,8 +899,12 @@ if __name__ == "__main__":
|
||||
default=False,
|
||||
help="Whether to use LoRA for finetuning.",
|
||||
)
|
||||
parser.add_argument("--lora_alpha", type=int, default=256, help="Alpha parameter for LoRA.")
|
||||
parser.add_argument("--lora_rank", type=int, default=128, help="LoRA rank parameter. ")
|
||||
parser.add_argument(
|
||||
"--lora_alpha", type=int, default=256, help="Alpha parameter for LoRA."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lora_rank", type=int, default=128, help="LoRA rank parameter. "
|
||||
)
|
||||
parser.add_argument("--fsdp_sharding_startegy", default="full")
|
||||
parser.add_argument("--multi_phased_distill_schedule", type=str, default=None)
|
||||
parser.add_argument(
|
||||
@@ -792,8 +919,10 @@ if __name__ == "__main__":
|
||||
"--lr_scheduler",
|
||||
type=str,
|
||||
default="constant",
|
||||
help=('The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
|
||||
' "constant", "constant_with_warmup"]'),
|
||||
help=(
|
||||
'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",'
|
||||
' "constant", "constant_with_warmup"]'
|
||||
),
|
||||
)
|
||||
parser.add_argument("--num_euler_timesteps", type=int, default=100)
|
||||
parser.add_argument(
|
||||
@@ -813,9 +942,13 @@ if __name__ == "__main__":
|
||||
action="store_true",
|
||||
help="Whether to apply the cfg_solver.",
|
||||
)
|
||||
parser.add_argument("--distill_cfg", type=float, default=3.0, help="Distillation coefficient.")
|
||||
parser.add_argument(
|
||||
"--distill_cfg", type=float, default=3.0, help="Distillation coefficient."
|
||||
)
|
||||
# ["euler_linear_quadratic", "pcm", "pcm_linear_qudratic"]
|
||||
parser.add_argument("--scheduler_type", type=str, default="pcm", help="The scheduler type to use.")
|
||||
parser.add_argument(
|
||||
"--scheduler_type", type=str, default="pcm", help="The scheduler type to use."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--adv_weight",
|
||||
type=float,
|
||||
@@ -828,18 +961,20 @@ if __name__ == "__main__":
|
||||
default=2,
|
||||
help="The stride of the discriminator head.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--linear_quadratic_threshold",
|
||||
type=float,
|
||||
default=0.025,
|
||||
help="The threshold of the linear quadratic scheduler.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--linear_range",
|
||||
type=float,
|
||||
default=0.5,
|
||||
help="Range for linear quadratic scheduler.",
|
||||
)
|
||||
parser.add_argument("--weight_decay", type=float, default=0.001, help="Weight decay to apply.")
|
||||
parser.add_argument(
|
||||
"--linear_quadratic_threshold",
|
||||
type=float,
|
||||
default=0.025,
|
||||
help="The threshold of the linear quadratic scheduler.",
|
||||
"--weight_decay", type=float, default=0.001, help="Weight decay to apply."
|
||||
)
|
||||
parser.add_argument(
|
||||
"--master_weight_type",
|
||||
@@ -848,4 +983,4 @@ if __name__ == "__main__":
|
||||
help="Weight type to use - fp32 or bf16.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
main(args)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,15 +1,19 @@
|
||||
from einops import rearrange
|
||||
from flash_attn import flash_attn_varlen_qkvpacked_func
|
||||
from flash_attn.bert_padding import pad_input, unpad_input
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
def flash_attn_no_pad(qkv, key_padding_mask, causal=False, dropout_p=0.0, softmax_scale=None):
|
||||
def flash_attn_no_pad(
|
||||
qkv, key_padding_mask, causal=False, dropout_p=0.0, softmax_scale=None
|
||||
):
|
||||
# adapted from https://github.com/Dao-AILab/flash-attention/blob/13403e81157ba37ca525890f2f0f2137edf75311/flash_attn/flash_attention.py#L27
|
||||
batch_size = qkv.shape[0]
|
||||
seqlen = qkv.shape[1]
|
||||
nheads = qkv.shape[-2]
|
||||
x = rearrange(qkv, "b s three h d -> b s (three h d)")
|
||||
x_unpad, indices, cu_seqlens, max_s, used_seqlens_in_batch = unpad_input(x, key_padding_mask)
|
||||
x_unpad, indices, cu_seqlens, max_s, used_seqlens_in_batch = unpad_input(
|
||||
x, key_padding_mask
|
||||
)
|
||||
|
||||
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=nheads)
|
||||
output_unpad = flash_attn_varlen_qkvpacked_func(
|
||||
@@ -21,7 +25,9 @@ def flash_attn_no_pad(qkv, key_padding_mask, causal=False, dropout_p=0.0, softma
|
||||
causal=causal,
|
||||
)
|
||||
output = rearrange(
|
||||
pad_input(rearrange(output_unpad, "nnz h d -> nnz (h d)"), indices, batch_size, seqlen),
|
||||
pad_input(
|
||||
rearrange(output_unpad, "nnz h d -> nnz (h d)"), indices, batch_size, seqlen
|
||||
),
|
||||
"b s (h d) -> b s h d",
|
||||
h=nheads,
|
||||
)
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
import os
|
||||
|
||||
import torch
|
||||
|
||||
__all__ = [
|
||||
@@ -34,7 +33,8 @@ C_SCALE = 1_000_000_000_000_000
|
||||
PROMPT_TEMPLATE_ENCODE = (
|
||||
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the image by detailing the color, shape, size, texture, "
|
||||
"quantity, text, spatial relationships of the objects and background:<|eot_id|>"
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
|
||||
)
|
||||
PROMPT_TEMPLATE_ENCODE_VIDEO = (
|
||||
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
|
||||
"1. The main content and theme of the video."
|
||||
@@ -42,15 +42,13 @@ PROMPT_TEMPLATE_ENCODE_VIDEO = (
|
||||
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
|
||||
"4. background environment, light, style and atmosphere."
|
||||
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
|
||||
)
|
||||
|
||||
NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
|
||||
|
||||
PROMPT_TEMPLATE = {
|
||||
"dit-llm-encode": {
|
||||
"template": PROMPT_TEMPLATE_ENCODE,
|
||||
"crop_start": 36,
|
||||
},
|
||||
"dit-llm-encode": {"template": PROMPT_TEMPLATE_ENCODE, "crop_start": 36,},
|
||||
"dit-llm-encode-video": {
|
||||
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
|
||||
"crop_start": 95,
|
||||
|
||||
@@ -1,3 +1,2 @@
|
||||
# ruff: noqa: F401
|
||||
from .pipelines import HunyuanVideoPipeline
|
||||
from .schedulers import FlowMatchDiscreteScheduler
|
||||
|
||||
@@ -1,2 +1 @@
|
||||
# ruff: noqa: F401
|
||||
from .pipeline_hunyuan_video import HunyuanVideoPipeline
|
||||
|
||||
@@ -17,33 +17,41 @@
|
||||
#
|
||||
# ==============================================================================
|
||||
import inspect
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
from typing import Any, Callable, Dict, List, Optional, Union, Tuple
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn.functional as F
|
||||
import numpy as np
|
||||
from dataclasses import dataclass
|
||||
from packaging import version
|
||||
|
||||
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
|
||||
from diffusers.configuration_utils import FrozenDict
|
||||
from diffusers.image_processor import VaeImageProcessor
|
||||
from diffusers.loaders import LoraLoaderMixin, TextualInversionLoaderMixin
|
||||
from diffusers.models import AutoencoderKL
|
||||
from diffusers.models.lora import adjust_lora_scale_text_encoder
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.schedulers import KarrasDiffusionSchedulers
|
||||
from diffusers.utils import (USE_PEFT_BACKEND, BaseOutput, deprecate, logging, replace_example_docstring,
|
||||
scale_lora_layers)
|
||||
from diffusers.utils import (
|
||||
USE_PEFT_BACKEND,
|
||||
deprecate,
|
||||
logging,
|
||||
replace_example_docstring,
|
||||
scale_lora_layers,
|
||||
unscale_lora_layers,
|
||||
)
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.utils.communications import all_gather
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.utils import BaseOutput
|
||||
|
||||
from ...constants import PRECISION_TO_TYPE
|
||||
from ...modules import HYVideoDiffusionTransformer
|
||||
from ...text_encoder import TextEncoder
|
||||
from ...vae.autoencoder_kl_causal_3d import AutoencoderKLCausal3D
|
||||
from ...text_encoder import TextEncoder
|
||||
from ...modules import HYVideoDiffusionTransformer
|
||||
|
||||
from einops import rearrange
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import all_gather, all_to_all_4D
|
||||
import torch.nn.functional as F
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
@@ -55,12 +63,16 @@ def rescale_noise_cfg(noise_cfg, noise_pred_text, guidance_rescale=0.0):
|
||||
Rescale `noise_cfg` according to `guidance_rescale`. Based on findings of [Common Diffusion Noise Schedules and
|
||||
Sample Steps are Flawed](https://arxiv.org/pdf/2305.08891.pdf). See Section 3.4
|
||||
"""
|
||||
std_text = noise_pred_text.std(dim=list(range(1, noise_pred_text.ndim)), keepdim=True)
|
||||
std_text = noise_pred_text.std(
|
||||
dim=list(range(1, noise_pred_text.ndim)), keepdim=True
|
||||
)
|
||||
std_cfg = noise_cfg.std(dim=list(range(1, noise_cfg.ndim)), keepdim=True)
|
||||
# rescale the results from guidance (fixes overexposure)
|
||||
noise_pred_rescaled = noise_cfg * (std_text / std_cfg)
|
||||
# mix with the original results from guidance by factor guidance_rescale to avoid "plain looking" images
|
||||
noise_cfg = (guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg)
|
||||
noise_cfg = (
|
||||
guidance_rescale * noise_pred_rescaled + (1 - guidance_rescale) * noise_cfg
|
||||
)
|
||||
return noise_cfg
|
||||
|
||||
|
||||
@@ -96,22 +108,30 @@ def retrieve_timesteps(
|
||||
second element is the number of inference steps.
|
||||
"""
|
||||
if timesteps is not None and sigmas is not None:
|
||||
raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
|
||||
raise ValueError(
|
||||
"Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values"
|
||||
)
|
||||
if timesteps is not None:
|
||||
accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
|
||||
accepts_timesteps = "timesteps" in set(
|
||||
inspect.signature(scheduler.set_timesteps).parameters.keys()
|
||||
)
|
||||
if not accepts_timesteps:
|
||||
raise ValueError(
|
||||
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
|
||||
f" timestep schedules. Please check whether you are using the correct scheduler.")
|
||||
f" timestep schedules. Please check whether you are using the correct scheduler."
|
||||
)
|
||||
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
elif sigmas is not None:
|
||||
accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
|
||||
accept_sigmas = "sigmas" in set(
|
||||
inspect.signature(scheduler.set_timesteps).parameters.keys()
|
||||
)
|
||||
if not accept_sigmas:
|
||||
raise ValueError(
|
||||
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
|
||||
f" sigmas schedules. Please check whether you are using the correct scheduler.")
|
||||
f" sigmas schedules. Please check whether you are using the correct scheduler."
|
||||
)
|
||||
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
@@ -173,27 +193,39 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
self.args = args
|
||||
# ==========================================================================================
|
||||
|
||||
if (hasattr(scheduler.config, "steps_offset") and scheduler.config.steps_offset != 1):
|
||||
if (
|
||||
hasattr(scheduler.config, "steps_offset")
|
||||
and scheduler.config.steps_offset != 1
|
||||
):
|
||||
deprecation_message = (
|
||||
f"The configuration file of this scheduler: {scheduler} is outdated. `steps_offset`"
|
||||
f" should be set to 1 instead of {scheduler.config.steps_offset}. Please make sure "
|
||||
"to update the config accordingly as leaving `steps_offset` might led to incorrect results"
|
||||
" in future versions. If you have downloaded this checkpoint from the Hugging Face Hub,"
|
||||
" it would be very nice if you could open a Pull request for the `scheduler/scheduler_config.json`"
|
||||
" file")
|
||||
deprecate("steps_offset!=1", "1.0.0", deprecation_message, standard_warn=False)
|
||||
" file"
|
||||
)
|
||||
deprecate(
|
||||
"steps_offset!=1", "1.0.0", deprecation_message, standard_warn=False
|
||||
)
|
||||
new_config = dict(scheduler.config)
|
||||
new_config["steps_offset"] = 1
|
||||
scheduler._internal_dict = FrozenDict(new_config)
|
||||
|
||||
if (hasattr(scheduler.config, "clip_sample") and scheduler.config.clip_sample is True):
|
||||
if (
|
||||
hasattr(scheduler.config, "clip_sample")
|
||||
and scheduler.config.clip_sample is True
|
||||
):
|
||||
deprecation_message = (
|
||||
f"The configuration file of this scheduler: {scheduler} has not set the configuration `clip_sample`."
|
||||
" `clip_sample` should be set to False in the configuration file. Please make sure to update the"
|
||||
" config accordingly as not setting `clip_sample` in the config might lead to incorrect results in"
|
||||
" future versions. If you have downloaded this checkpoint from the Hugging Face Hub, it would be very"
|
||||
" nice if you could open a Pull request for the `scheduler/scheduler_config.json` file")
|
||||
deprecate("clip_sample not set", "1.0.0", deprecation_message, standard_warn=False)
|
||||
" nice if you could open a Pull request for the `scheduler/scheduler_config.json` file"
|
||||
)
|
||||
deprecate(
|
||||
"clip_sample not set", "1.0.0", deprecation_message, standard_warn=False
|
||||
)
|
||||
new_config = dict(scheduler.config)
|
||||
new_config["clip_sample"] = False
|
||||
scheduler._internal_dict = FrozenDict(new_config)
|
||||
@@ -205,7 +237,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
scheduler=scheduler,
|
||||
text_encoder_2=text_encoder_2,
|
||||
)
|
||||
self.vae_scale_factor = 2**(len(self.vae.config.block_out_channels) - 1)
|
||||
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
|
||||
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
|
||||
|
||||
def encode_prompt(
|
||||
@@ -271,6 +303,13 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
else:
|
||||
scale_lora_layers(text_encoder.model, lora_scale)
|
||||
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
if prompt_embeds is None:
|
||||
# textual inversion: process multi-vector tokens if necessary
|
||||
if isinstance(self, TextualInversionLoaderMixin):
|
||||
@@ -278,7 +317,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
|
||||
text_inputs = text_encoder.text2tokens(prompt, data_type=data_type)
|
||||
if clip_skip is None:
|
||||
prompt_outputs = text_encoder.encode(text_inputs, data_type=data_type, device=device)
|
||||
prompt_outputs = text_encoder.encode(
|
||||
text_inputs, data_type=data_type, device=device
|
||||
)
|
||||
prompt_embeds = prompt_outputs.hidden_state
|
||||
else:
|
||||
prompt_outputs = text_encoder.encode(
|
||||
@@ -295,14 +336,18 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
# representations. The `last_hidden_states` that we typically use for
|
||||
# obtaining the final prompt representations passes through the LayerNorm
|
||||
# layer.
|
||||
prompt_embeds = text_encoder.model.text_model.final_layer_norm(prompt_embeds)
|
||||
prompt_embeds = text_encoder.model.text_model.final_layer_norm(
|
||||
prompt_embeds
|
||||
)
|
||||
|
||||
attention_mask = prompt_outputs.attention_mask
|
||||
if attention_mask is not None:
|
||||
attention_mask = attention_mask.to(device)
|
||||
bs_embed, seq_len = attention_mask.shape
|
||||
attention_mask = attention_mask.repeat(1, num_videos_per_prompt)
|
||||
attention_mask = attention_mask.view(bs_embed * num_videos_per_prompt, seq_len)
|
||||
attention_mask = attention_mask.view(
|
||||
bs_embed * num_videos_per_prompt, seq_len
|
||||
)
|
||||
|
||||
if text_encoder is not None:
|
||||
prompt_embeds_dtype = text_encoder.dtype
|
||||
@@ -322,7 +367,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
bs_embed, seq_len, _ = prompt_embeds.shape
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
|
||||
prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, seq_len, -1)
|
||||
prompt_embeds = prompt_embeds.view(
|
||||
bs_embed * num_videos_per_prompt, seq_len, -1
|
||||
)
|
||||
|
||||
return (
|
||||
prompt_embeds,
|
||||
@@ -338,7 +385,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
latents = 1 / self.vae.config.scaling_factor * latents
|
||||
if enable_tiling:
|
||||
self.vae.enable_tiling()
|
||||
image = self.vae.decode(latents, return_dict=False)[0]
|
||||
image = self.vae.decode(latents, return_dict=False)[0]
|
||||
else:
|
||||
image = self.vae.decode(latents, return_dict=False)[0]
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloat16
|
||||
if image.ndim == 4:
|
||||
@@ -374,21 +423,33 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
vae_ver="88-4c-sd",
|
||||
):
|
||||
if height % 8 != 0 or width % 8 != 0:
|
||||
raise ValueError(f"`height` and `width` have to be divisible by 8 but are {height} and {width}.")
|
||||
raise ValueError(
|
||||
f"`height` and `width` have to be divisible by 8 but are {height} and {width}."
|
||||
)
|
||||
|
||||
if video_length is not None:
|
||||
if "884" in vae_ver:
|
||||
if video_length != 1 and (video_length - 1) % 4 != 0:
|
||||
raise ValueError(f"`video_length` has to be 1 or a multiple of 4 but is {video_length}.")
|
||||
raise ValueError(
|
||||
f"`video_length` has to be 1 or a multiple of 4 but is {video_length}."
|
||||
)
|
||||
elif "888" in vae_ver:
|
||||
if video_length != 1 and (video_length - 1) % 8 != 0:
|
||||
raise ValueError(f"`video_length` has to be 1 or a multiple of 8 but is {video_length}.")
|
||||
raise ValueError(
|
||||
f"`video_length` has to be 1 or a multiple of 8 but is {video_length}."
|
||||
)
|
||||
|
||||
if callback_steps is not None and (not isinstance(callback_steps, int) or callback_steps <= 0):
|
||||
raise ValueError(f"`callback_steps` has to be a positive integer but is {callback_steps} of type"
|
||||
f" {type(callback_steps)}.")
|
||||
if callback_on_step_end_tensor_inputs is not None and not all(k in self._callback_tensor_inputs
|
||||
for k in callback_on_step_end_tensor_inputs):
|
||||
if callback_steps is not None and (
|
||||
not isinstance(callback_steps, int) or callback_steps <= 0
|
||||
):
|
||||
raise ValueError(
|
||||
f"`callback_steps` has to be a positive integer but is {callback_steps} of type"
|
||||
f" {type(callback_steps)}."
|
||||
)
|
||||
if callback_on_step_end_tensor_inputs is not None and not all(
|
||||
k in self._callback_tensor_inputs
|
||||
for k in callback_on_step_end_tensor_inputs
|
||||
):
|
||||
raise ValueError(
|
||||
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
|
||||
)
|
||||
@@ -396,23 +457,32 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
if prompt is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two.")
|
||||
" only forward one of the two."
|
||||
)
|
||||
elif prompt is None and prompt_embeds is None:
|
||||
raise ValueError(
|
||||
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined.")
|
||||
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
|
||||
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
|
||||
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
|
||||
)
|
||||
elif prompt is not None and (
|
||||
not isinstance(prompt, str) and not isinstance(prompt, list)
|
||||
):
|
||||
raise ValueError(
|
||||
f"`prompt` has to be of type `str` or `list` but is {type(prompt)}"
|
||||
)
|
||||
|
||||
if negative_prompt is not None and negative_prompt_embeds is not None:
|
||||
raise ValueError(f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
|
||||
f" {negative_prompt_embeds}. Please make sure to only forward one of the two.")
|
||||
raise ValueError(
|
||||
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
|
||||
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
|
||||
)
|
||||
|
||||
if prompt_embeds is not None and negative_prompt_embeds is not None:
|
||||
if prompt_embeds.shape != negative_prompt_embeds.shape:
|
||||
raise ValueError(
|
||||
"`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but"
|
||||
f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`"
|
||||
f" {negative_prompt_embeds.shape}.")
|
||||
f" {negative_prompt_embeds.shape}."
|
||||
)
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
@@ -436,10 +506,13 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators.")
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
|
||||
if latents is None:
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
latents = randn_tensor(
|
||||
shape, generator=generator, device=device, dtype=dtype
|
||||
)
|
||||
else:
|
||||
latents = latents.to(device)
|
||||
|
||||
@@ -542,15 +615,18 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
guidance_rescale: float = 0.0,
|
||||
clip_skip: Optional[int] = None,
|
||||
callback_on_step_end: Optional[Union[Callable[[int, int, Dict], None], PipelineCallback,
|
||||
MultiPipelineCallbacks, ]] = None,
|
||||
callback_on_step_end: Optional[
|
||||
Union[
|
||||
Callable[[int, int, Dict], None],
|
||||
PipelineCallback,
|
||||
MultiPipelineCallbacks,
|
||||
]
|
||||
] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
vae_ver: str = "88-4c-sd",
|
||||
enable_tiling: bool = False,
|
||||
enable_vae_sp: bool = False,
|
||||
n_tokens: Optional[int] = None,
|
||||
embedded_guidance_scale: Optional[float] = None,
|
||||
mask_strategy: Optional[Dict[str, list]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
r"""
|
||||
@@ -687,11 +763,18 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
device = (torch.device(f"cuda:{dist.get_rank()}") if dist.is_initialized() else self._execution_device)
|
||||
device = (
|
||||
torch.device(f"cuda:{dist.get_rank()}")
|
||||
if dist.is_initialized()
|
||||
else self._execution_device
|
||||
)
|
||||
|
||||
# 3. Encode input prompt
|
||||
lora_scale = (self.cross_attention_kwargs.get("scale", None)
|
||||
if self.cross_attention_kwargs is not None else None)
|
||||
lora_scale = (
|
||||
self.cross_attention_kwargs.get("scale", None)
|
||||
if self.cross_attention_kwargs is not None
|
||||
else None
|
||||
)
|
||||
|
||||
(
|
||||
prompt_embeds,
|
||||
@@ -752,8 +835,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
prompt_mask_2 = torch.cat([negative_prompt_mask_2, prompt_mask_2])
|
||||
|
||||
# 4. Prepare timesteps
|
||||
extra_set_timesteps_kwargs = self.prepare_extra_func_kwargs(self.scheduler.set_timesteps,
|
||||
{"n_tokens": n_tokens})
|
||||
extra_set_timesteps_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.scheduler.set_timesteps, {"n_tokens": n_tokens}
|
||||
)
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
@@ -782,39 +866,32 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
|
||||
if get_sequence_parallel_state():
|
||||
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
|
||||
latents = rearrange(
|
||||
latents, "b t (n s) h w -> b t n s h w", n=world_size
|
||||
).contiguous()
|
||||
latents = latents[:, :, rank, :, :, :]
|
||||
|
||||
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
|
||||
extra_step_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.scheduler.step,
|
||||
{
|
||||
"generator": generator,
|
||||
"eta": eta
|
||||
},
|
||||
self.scheduler.step, {"generator": generator, "eta": eta},
|
||||
)
|
||||
|
||||
target_dtype = PRECISION_TO_TYPE[self.args.precision]
|
||||
autocast_enabled = (target_dtype != torch.float32) and not self.args.disable_autocast
|
||||
autocast_enabled = (
|
||||
target_dtype != torch.float32
|
||||
) and not self.args.disable_autocast
|
||||
vae_dtype = PRECISION_TO_TYPE[self.args.vae_precision]
|
||||
vae_autocast_enabled = (vae_dtype != torch.float32) and not self.args.disable_autocast
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32
|
||||
) and not self.args.disable_autocast
|
||||
|
||||
# 7. Denoising loop
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
def dict_to_3d_list(mask_strategy, t_max=50, l_max=60, h_max=24):
|
||||
result = [[[None for _ in range(h_max)] for _ in range(l_max)] for _ in range(t_max)]
|
||||
if mask_strategy is None:
|
||||
return result
|
||||
for key, value in mask_strategy.items():
|
||||
t, l, h = map(int, key.split('_'))
|
||||
result[t][l][h] = value
|
||||
return result
|
||||
|
||||
mask_strategy = dict_to_3d_list(mask_strategy)
|
||||
# if is_progress_bar:
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
@@ -822,39 +899,57 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
continue
|
||||
|
||||
# expand the latents if we are doing classifier free guidance
|
||||
latent_model_input = (torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents)
|
||||
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
|
||||
latent_model_input = (
|
||||
torch.cat([latents] * 2)
|
||||
if self.do_classifier_free_guidance
|
||||
else latents
|
||||
)
|
||||
latent_model_input = self.scheduler.scale_model_input(
|
||||
latent_model_input, t
|
||||
)
|
||||
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
guidance_expand = (torch.tensor(
|
||||
[embedded_guidance_scale] * latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
).to(target_dtype) * 1000.0 if embedded_guidance_scale is not None else None)
|
||||
guidance_expand = (
|
||||
torch.tensor(
|
||||
[embedded_guidance_scale] * latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=device,
|
||||
).to(target_dtype)
|
||||
* 1000.0
|
||||
if embedded_guidance_scale is not None
|
||||
else None
|
||||
)
|
||||
# predict the noise residual
|
||||
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
|
||||
# concat prompt_embeds_2 and prompt_embeds. Mismatch fill with zeros
|
||||
with torch.autocast(
|
||||
device_type="cuda", dtype=target_dtype, enabled=autocast_enabled
|
||||
):
|
||||
# concat prompt_embeds_2 and prompt_embeds. Mismach fill with zeros
|
||||
if prompt_embeds_2.shape[-1] != prompt_embeds.shape[-1]:
|
||||
prompt_embeds_2 = F.pad(
|
||||
prompt_embeds_2,
|
||||
(0, prompt_embeds.shape[2] - prompt_embeds_2.shape[1]),
|
||||
value=0,
|
||||
).unsqueeze(1)
|
||||
encoder_hidden_states = torch.cat([prompt_embeds_2, prompt_embeds], dim=1)
|
||||
encoder_hidden_states = torch.cat(
|
||||
[prompt_embeds_2, prompt_embeds], dim=1
|
||||
)
|
||||
noise_pred = self.transformer( # For an input image (129, 192, 336) (1, 256, 256)
|
||||
latent_model_input,
|
||||
latent_model_input, # [2, 16, 33, 24, 42]
|
||||
encoder_hidden_states,
|
||||
t_expand,
|
||||
prompt_mask,
|
||||
mask_strategy=mask_strategy[i],
|
||||
t_expand, # [2]
|
||||
prompt_mask, # [2, 256]fpdb
|
||||
guidance=guidance_expand,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
)[
|
||||
0
|
||||
]
|
||||
|
||||
# perform guidance
|
||||
if self.do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
noise_pred = noise_pred_uncond + self.guidance_scale * (
|
||||
noise_pred_text - noise_pred_uncond
|
||||
)
|
||||
|
||||
if self.do_classifier_free_guidance and self.guidance_rescale > 0.0:
|
||||
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
|
||||
@@ -865,7 +960,9 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, **extra_step_kwargs, return_dict=False)[0]
|
||||
latents = self.scheduler.step(
|
||||
noise_pred, t, latents, **extra_step_kwargs, return_dict=False
|
||||
)[0]
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
@@ -875,10 +972,14 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
|
||||
negative_prompt_embeds = callback_outputs.pop(
|
||||
"negative_prompt_embeds", negative_prompt_embeds
|
||||
)
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
if i == len(timesteps) - 1 or (
|
||||
(i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0
|
||||
):
|
||||
if progress_bar is not None:
|
||||
progress_bar.update()
|
||||
if callback is not None and i % callback_steps == 0:
|
||||
@@ -898,19 +999,32 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
pass
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Only support latents with shape (b, c, h, w) or (b, c, f, h, w), but got {latents.shape}.")
|
||||
f"Only support latents with shape (b, c, h, w) or (b, c, f, h, w), but got {latents.shape}."
|
||||
)
|
||||
|
||||
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 (
|
||||
hasattr(self.vae.config, "shift_factor")
|
||||
and self.vae.config.shift_factor
|
||||
):
|
||||
latents = (
|
||||
latents / self.vae.config.scaling_factor
|
||||
+ self.vae.config.shift_factor
|
||||
)
|
||||
else:
|
||||
latents = latents / self.vae.config.scaling_factor
|
||||
|
||||
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
|
||||
with torch.autocast(
|
||||
device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled
|
||||
):
|
||||
if enable_tiling:
|
||||
self.vae.enable_tiling()
|
||||
if enable_vae_sp:
|
||||
self.vae.enable_parallel()
|
||||
image = self.vae.decode(latents, return_dict=False, generator=generator)[0]
|
||||
image = self.vae.decode(
|
||||
latents, return_dict=False, generator=generator
|
||||
)[0]
|
||||
else:
|
||||
image = self.vae.decode(
|
||||
latents, return_dict=False, generator=generator
|
||||
)[0]
|
||||
|
||||
if expand_temporal_dim or image.shape[2] == 1:
|
||||
image = image.squeeze(2)
|
||||
|
||||
@@ -1,2 +1 @@
|
||||
# ruff: noqa: F401
|
||||
from .scheduling_flow_match_discrete import FlowMatchDiscreteScheduler
|
||||
|
||||
@@ -20,10 +20,13 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
@@ -87,7 +90,9 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
|
||||
self.supported_solver = ["euler"]
|
||||
if solver not in self.supported_solver:
|
||||
raise ValueError(f"Solver {solver} not supported. Supported solvers: {self.supported_solver}")
|
||||
raise ValueError(
|
||||
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
|
||||
)
|
||||
|
||||
@property
|
||||
def step_index(self):
|
||||
@@ -143,7 +148,9 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
sigmas = 1 - sigmas
|
||||
|
||||
self.sigmas = sigmas
|
||||
self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(dtype=torch.float32, device=device)
|
||||
self.timesteps = (sigmas[:-1] * self.config.num_train_timesteps).to(
|
||||
dtype=torch.float32, device=device
|
||||
)
|
||||
|
||||
# Reset step index
|
||||
self._step_index = None
|
||||
@@ -170,7 +177,9 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
else:
|
||||
self._step_index = self._begin_index
|
||||
|
||||
def scale_model_input(self, sample: torch.Tensor, timestep: Optional[int] = None) -> torch.Tensor:
|
||||
def scale_model_input(
|
||||
self, sample: torch.Tensor, timestep: Optional[int] = None
|
||||
) -> torch.Tensor:
|
||||
return sample
|
||||
|
||||
def sd3_time_shift(self, t: torch.Tensor):
|
||||
@@ -208,11 +217,18 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
returned, otherwise a tuple is returned where the first element is the sample tensor.
|
||||
"""
|
||||
|
||||
if (isinstance(timestep, int) or isinstance(timestep, torch.IntTensor)
|
||||
or isinstance(timestep, torch.LongTensor)):
|
||||
raise ValueError(("Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||
" one of the `scheduler.timesteps` as a timestep."), )
|
||||
if (
|
||||
isinstance(timestep, int)
|
||||
or isinstance(timestep, torch.IntTensor)
|
||||
or isinstance(timestep, torch.LongTensor)
|
||||
):
|
||||
raise ValueError(
|
||||
(
|
||||
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
|
||||
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
|
||||
" one of the `scheduler.timesteps` as a timestep."
|
||||
),
|
||||
)
|
||||
|
||||
if self.step_index is None:
|
||||
self._init_step_index(timestep)
|
||||
@@ -225,13 +241,15 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
if self.config.solver == "euler":
|
||||
prev_sample = sample + model_output.to(torch.float32) * dt
|
||||
else:
|
||||
raise ValueError(f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}")
|
||||
raise ValueError(
|
||||
f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}"
|
||||
)
|
||||
|
||||
# upon completion increase step index by one
|
||||
self._step_index += 1
|
||||
|
||||
if not return_dict:
|
||||
return (prev_sample, )
|
||||
return (prev_sample,)
|
||||
|
||||
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
# ruff: noqa: F405, F403
|
||||
import argparse
|
||||
import re
|
||||
|
||||
from .constants import *
|
||||
import re
|
||||
from .modules.models import HUNYUAN_VIDEO_CONFIG
|
||||
|
||||
|
||||
@@ -47,12 +45,16 @@ def add_network_args(parser: argparse.ArgumentParser):
|
||||
)
|
||||
|
||||
# RoPE
|
||||
group.add_argument("--rope-theta", type=int, default=256, help="Theta used in RoPE.")
|
||||
group.add_argument(
|
||||
"--rope-theta", type=int, default=256, help="Theta used in RoPE."
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
def add_extra_models_args(parser: argparse.ArgumentParser):
|
||||
group = parser.add_argument_group(title="Extra models args, including vae, text encoders and tokenizers)")
|
||||
group = parser.add_argument_group(
|
||||
title="Extra models args, including vae, text encoders and tokenizers)"
|
||||
)
|
||||
|
||||
# - VAE
|
||||
group.add_argument(
|
||||
@@ -96,7 +98,9 @@ def add_extra_models_args(parser: argparse.ArgumentParser):
|
||||
default=4096,
|
||||
help="Dimension of the text encoder hidden states.",
|
||||
)
|
||||
group.add_argument("--text-len", type=int, default=256, help="Maximum length of the text input.")
|
||||
group.add_argument(
|
||||
"--text-len", type=int, default=256, help="Maximum length of the text input."
|
||||
)
|
||||
group.add_argument(
|
||||
"--tokenizer",
|
||||
type=str,
|
||||
@@ -191,10 +195,7 @@ def add_denoise_schedule_args(parser: argparse.ArgumentParser):
|
||||
help="If reverse, learning/sampling from t=1 -> t=0.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--flow-solver",
|
||||
type=str,
|
||||
default="euler",
|
||||
help="Solver for flow matching.",
|
||||
"--flow-solver", type=str, default="euler", help="Solver for flow matching.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--use-linear-quadratic-schedule",
|
||||
@@ -329,13 +330,17 @@ def add_inference_args(parser: argparse.ArgumentParser):
|
||||
group.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
|
||||
|
||||
# Classifier-Free Guidance
|
||||
group.add_argument("--neg-prompt", type=str, default=None, help="Negative prompt for sampling.")
|
||||
group.add_argument("--cfg-scale", type=float, default=1.0, help="Classifier free guidance scale.")
|
||||
group.add_argument(
|
||||
"--neg-prompt", type=str, default=None, help="Negative prompt for sampling."
|
||||
)
|
||||
group.add_argument(
|
||||
"--cfg-scale", type=float, default=1.0, help="Classifier free guidance scale."
|
||||
)
|
||||
group.add_argument(
|
||||
"--embedded-cfg-scale",
|
||||
type=float,
|
||||
default=6.0,
|
||||
help="Embedded classifier free guidance scale.",
|
||||
help="Embeded classifier free guidance scale.",
|
||||
)
|
||||
|
||||
group.add_argument(
|
||||
@@ -352,16 +357,10 @@ def add_parallel_args(parser: argparse.ArgumentParser):
|
||||
|
||||
# ======================== Model loads ========================
|
||||
group.add_argument(
|
||||
"--ulysses-degree",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Ulysses degree.",
|
||||
"--ulysses-degree", type=int, default=1, help="Ulysses degree.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--ring-degree",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Ulysses degree.",
|
||||
"--ring-degree", type=int, default=1, help="Ulysses degree.",
|
||||
)
|
||||
|
||||
return parser
|
||||
@@ -371,10 +370,14 @@ def sanity_check_args(args):
|
||||
# VAE channels
|
||||
vae_pattern = r"\d{2,3}-\d{1,2}c-\w+"
|
||||
if not re.match(vae_pattern, args.vae):
|
||||
raise ValueError(f"Invalid VAE model: {args.vae}. Must be in the format of '{vae_pattern}'.")
|
||||
raise ValueError(
|
||||
f"Invalid VAE model: {args.vae}. Must be in the format of '{vae_pattern}'."
|
||||
)
|
||||
vae_channels = int(args.vae.split("-")[1][:-1])
|
||||
if args.latent_channels is None:
|
||||
args.latent_channels = vae_channels
|
||||
if vae_channels != args.latent_channels:
|
||||
raise ValueError(f"Latent channels ({args.latent_channels}) must match the VAE channels ({vae_channels}).")
|
||||
raise ValueError(
|
||||
f"Latent channels ({args.latent_channels}) must match the VAE channels ({vae_channels})."
|
||||
)
|
||||
return args
|
||||
|
||||
@@ -1,24 +1,33 @@
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
import random
|
||||
import functools
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
from pathlib import Path
|
||||
from loguru import logger
|
||||
|
||||
import torch
|
||||
from loguru import logger
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
|
||||
from fastvideo.models.hunyuan.constants import NEGATIVE_PROMPT, PRECISION_TO_TYPE, PROMPT_TEMPLATE
|
||||
from fastvideo.models.hunyuan.diffusion.pipelines import HunyuanVideoPipeline
|
||||
from fastvideo.models.hunyuan.diffusion.schedulers import FlowMatchDiscreteScheduler
|
||||
import torch.distributed as dist
|
||||
from fastvideo.models.hunyuan.constants import (
|
||||
PROMPT_TEMPLATE,
|
||||
NEGATIVE_PROMPT,
|
||||
PRECISION_TO_TYPE,
|
||||
)
|
||||
from fastvideo.models.hunyuan.vae import load_vae
|
||||
from fastvideo.models.hunyuan.modules import load_model
|
||||
from fastvideo.models.hunyuan.text_encoder import TextEncoder
|
||||
from fastvideo.models.hunyuan.utils.data_utils import align_to
|
||||
from fastvideo.models.hunyuan.vae import load_vae
|
||||
from fastvideo.utils.parallel_states import nccl_info
|
||||
from fastvideo.models.hunyuan.diffusion.schedulers import FlowMatchDiscreteScheduler
|
||||
from fastvideo.models.hunyuan.diffusion.pipelines import HunyuanVideoPipeline
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
from fastvideo.utils.parallel_states import (
|
||||
initialize_sequence_parallel_state,
|
||||
nccl_info,
|
||||
)
|
||||
|
||||
|
||||
class Inference(object):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
args,
|
||||
@@ -44,7 +53,13 @@ class Inference(object):
|
||||
self.use_cpu_offload = use_cpu_offload
|
||||
|
||||
self.args = args
|
||||
self.device = (device if device is not None else "cuda" if torch.cuda.is_available() else "cpu")
|
||||
self.device = (
|
||||
device
|
||||
if device is not None
|
||||
else "cuda"
|
||||
if torch.cuda.is_available()
|
||||
else "cpu"
|
||||
)
|
||||
self.logger = logger
|
||||
self.parallel_args = parallel_args
|
||||
|
||||
@@ -88,8 +103,6 @@ class Inference(object):
|
||||
)
|
||||
model = model.to(device)
|
||||
model = Inference.load_state_dict(args, model, pretrained_model_path)
|
||||
if args.enable_torch_compile:
|
||||
model = torch.compile(model)
|
||||
model.eval()
|
||||
|
||||
# ============================= Build extra models ========================
|
||||
@@ -104,7 +117,9 @@ class Inference(object):
|
||||
|
||||
# Text encoder
|
||||
if args.prompt_template_video is not None:
|
||||
crop_start = PROMPT_TEMPLATE[args.prompt_template_video].get("crop_start", 0)
|
||||
crop_start = PROMPT_TEMPLATE[args.prompt_template_video].get(
|
||||
"crop_start", 0
|
||||
)
|
||||
elif args.prompt_template is not None:
|
||||
crop_start = PROMPT_TEMPLATE[args.prompt_template].get("crop_start", 0)
|
||||
else:
|
||||
@@ -112,11 +127,18 @@ class Inference(object):
|
||||
max_length = args.text_len + crop_start
|
||||
|
||||
# prompt_template
|
||||
prompt_template = (PROMPT_TEMPLATE[args.prompt_template] if args.prompt_template is not None else None)
|
||||
prompt_template = (
|
||||
PROMPT_TEMPLATE[args.prompt_template]
|
||||
if args.prompt_template is not None
|
||||
else None
|
||||
)
|
||||
|
||||
# prompt_template_video
|
||||
prompt_template_video = (PROMPT_TEMPLATE[args.prompt_template_video]
|
||||
if args.prompt_template_video is not None else None)
|
||||
prompt_template_video = (
|
||||
PROMPT_TEMPLATE[args.prompt_template_video]
|
||||
if args.prompt_template_video is not None
|
||||
else None
|
||||
)
|
||||
|
||||
text_encoder = TextEncoder(
|
||||
text_encoder_type=args.text_encoder,
|
||||
@@ -173,14 +195,18 @@ class Inference(object):
|
||||
files = [f for f in files if str(f).endswith("_model_states.pt")]
|
||||
model_path = files[0]
|
||||
if len(files) > 1:
|
||||
logger.warning(f"Multiple model weights found in {dit_weight}, using {model_path}")
|
||||
logger.warning(
|
||||
f"Multiple model weights found in {dit_weight}, using {model_path}"
|
||||
)
|
||||
bare_model = False
|
||||
else:
|
||||
raise ValueError(f"Invalid model path: {dit_weight} with unrecognized weight format: "
|
||||
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
|
||||
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
|
||||
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
|
||||
f"specific weight file, please provide the full path to the file.")
|
||||
raise ValueError(
|
||||
f"Invalid model path: {dit_weight} with unrecognized weight format: "
|
||||
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
|
||||
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
|
||||
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
|
||||
f"specific weight file, please provide the full path to the file."
|
||||
)
|
||||
else:
|
||||
if dit_weight.is_dir():
|
||||
files = list(dit_weight.glob("*.pt"))
|
||||
@@ -193,14 +219,18 @@ class Inference(object):
|
||||
files = [f for f in files if str(f).endswith("_model_states.pt")]
|
||||
model_path = files[0]
|
||||
if len(files) > 1:
|
||||
logger.warning(f"Multiple model weights found in {dit_weight}, using {model_path}")
|
||||
logger.warning(
|
||||
f"Multiple model weights found in {dit_weight}, using {model_path}"
|
||||
)
|
||||
bare_model = False
|
||||
else:
|
||||
raise ValueError(f"Invalid model path: {dit_weight} with unrecognized weight format: "
|
||||
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
|
||||
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
|
||||
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
|
||||
f"specific weight file, please provide the full path to the file.")
|
||||
raise ValueError(
|
||||
f"Invalid model path: {dit_weight} with unrecognized weight format: "
|
||||
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
|
||||
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
|
||||
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
|
||||
f"specific weight file, please provide the full path to the file."
|
||||
)
|
||||
elif dit_weight.is_file():
|
||||
model_path = dit_weight
|
||||
bare_model = "unknown"
|
||||
@@ -215,7 +245,9 @@ class Inference(object):
|
||||
state_dict = safetensors_load_file(model_path)
|
||||
elif model_path.suffix == ".pt":
|
||||
# Use torch for .pt files
|
||||
state_dict = torch.load(model_path, map_location=lambda storage, loc: storage)
|
||||
state_dict = torch.load(
|
||||
model_path, map_location=lambda storage, loc: storage
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported file format: {model_path}")
|
||||
|
||||
@@ -225,8 +257,10 @@ class Inference(object):
|
||||
if load_key in state_dict:
|
||||
state_dict = state_dict[load_key]
|
||||
else:
|
||||
raise KeyError(f"Missing key: `{load_key}` in the checkpoint: {model_path}. The keys in the checkpoint "
|
||||
f"are: {list(state_dict.keys())}.")
|
||||
raise KeyError(
|
||||
f"Missing key: `{load_key}` in the checkpoint: {model_path}. The keys in the checkpoint "
|
||||
f"are: {list(state_dict.keys())}."
|
||||
)
|
||||
model.load_state_dict(state_dict, strict=True)
|
||||
return model
|
||||
|
||||
@@ -244,7 +278,6 @@ class Inference(object):
|
||||
|
||||
|
||||
class HunyuanVideoSampler(Inference):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
args,
|
||||
@@ -338,7 +371,6 @@ class HunyuanVideoSampler(Inference):
|
||||
embedded_guidance_scale=None,
|
||||
batch_size=1,
|
||||
num_videos_per_prompt=1,
|
||||
mask_strategy=None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
@@ -365,20 +397,34 @@ class HunyuanVideoSampler(Inference):
|
||||
if isinstance(seed, torch.Tensor):
|
||||
seed = seed.tolist()
|
||||
if seed is None:
|
||||
seeds = [random.randint(0, 1_000_000) for _ in range(batch_size * num_videos_per_prompt)]
|
||||
seeds = [
|
||||
random.randint(0, 1_000_000)
|
||||
for _ in range(batch_size * num_videos_per_prompt)
|
||||
]
|
||||
elif isinstance(seed, int):
|
||||
seeds = [seed + i for _ in range(batch_size) for i in range(num_videos_per_prompt)]
|
||||
seeds = [
|
||||
seed + i
|
||||
for _ in range(batch_size)
|
||||
for i in range(num_videos_per_prompt)
|
||||
]
|
||||
elif isinstance(seed, (list, tuple)):
|
||||
if len(seed) == batch_size:
|
||||
seeds = [int(seed[i]) + j for i in range(batch_size) for j in range(num_videos_per_prompt)]
|
||||
seeds = [
|
||||
int(seed[i]) + j
|
||||
for i in range(batch_size)
|
||||
for j in range(num_videos_per_prompt)
|
||||
]
|
||||
elif len(seed) == batch_size * num_videos_per_prompt:
|
||||
seeds = [int(s) for s in seed]
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Length of seed must be equal to number of prompt(batch_size) or "
|
||||
f"batch_size * num_videos_per_prompt ({batch_size} * {num_videos_per_prompt}), got {seed}.")
|
||||
f"batch_size * num_videos_per_prompt ({batch_size} * {num_videos_per_prompt}), got {seed}."
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Seed must be an integer, a list of integers, or None, got {seed}.")
|
||||
raise ValueError(
|
||||
f"Seed must be an integer, a list of integers, or None, got {seed}."
|
||||
)
|
||||
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
|
||||
generator = [torch.Generator("cpu").manual_seed(seed) for seed in seeds]
|
||||
out_dict["seeds"] = seeds
|
||||
@@ -391,9 +437,13 @@ class HunyuanVideoSampler(Inference):
|
||||
f"`height` and `width` and `video_length` must be positive integers, got height={height}, width={width}, video_length={video_length}"
|
||||
)
|
||||
if (video_length - 1) % 4 != 0:
|
||||
raise ValueError(f"`video_length-1` must be a multiple of 4, got {video_length}")
|
||||
raise ValueError(
|
||||
f"`video_length-1` must be a multiple of 4, got {video_length}"
|
||||
)
|
||||
|
||||
logger.info(f"Input (height, width, video_length) = ({height}, {width}, {video_length})")
|
||||
logger.info(
|
||||
f"Input (height, width, video_length) = ({height}, {width}, {video_length})"
|
||||
)
|
||||
|
||||
target_height = align_to(height, 16)
|
||||
target_width = align_to(width, 16)
|
||||
@@ -412,7 +462,9 @@ class HunyuanVideoSampler(Inference):
|
||||
if negative_prompt is None or negative_prompt == "":
|
||||
negative_prompt = self.default_negative_prompt
|
||||
if not isinstance(negative_prompt, str):
|
||||
raise TypeError(f"`negative_prompt` must be a string, but got {type(negative_prompt)}")
|
||||
raise TypeError(
|
||||
f"`negative_prompt` must be a string, but got {type(negative_prompt)}"
|
||||
)
|
||||
negative_prompt = [negative_prompt.strip()]
|
||||
|
||||
# ========================================================================
|
||||
@@ -470,8 +522,6 @@ class HunyuanVideoSampler(Inference):
|
||||
is_progress_bar=True,
|
||||
vae_ver=self.args.vae,
|
||||
enable_tiling=self.args.vae_tiling,
|
||||
enable_vae_sp=self.args.vae_sp,
|
||||
mask_strategy=mask_strategy,
|
||||
)[0]
|
||||
out_dict["samples"] = samples
|
||||
out_dict["prompts"] = prompt
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from .models import HUNYUAN_VIDEO_CONFIG, HYVideoDiffusionTransformer
|
||||
from .models import HYVideoDiffusionTransformer, HUNYUAN_VIDEO_CONFIG
|
||||
|
||||
|
||||
def load_model(args, in_channels, out_channels, factor_kwargs):
|
||||
|
||||
@@ -1,25 +1,18 @@
|
||||
import importlib.metadata
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from einops import rearrange
|
||||
|
||||
try:
|
||||
from st_attn import sliding_tile_attention
|
||||
except ImportError:
|
||||
print("Could not load Sliding Tile Attention.")
|
||||
sliding_tile_attention = None
|
||||
|
||||
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
|
||||
from fastvideo.utils.communications import all_gather, all_to_all_4D
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
from fastvideo.utils.communications import all_gather, all_to_all_4D
|
||||
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
|
||||
|
||||
|
||||
def attention(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
drop_rate=0,
|
||||
attn_mask=None,
|
||||
causal=False,
|
||||
q, k, v, drop_rate=0, attn_mask=None, causal=False,
|
||||
):
|
||||
|
||||
qkv = torch.stack([q, k, v], dim=2)
|
||||
@@ -27,43 +20,21 @@ def attention(
|
||||
if attn_mask is not None and attn_mask.dtype != torch.bool:
|
||||
attn_mask = attn_mask.bool()
|
||||
|
||||
x = flash_attn_no_pad(qkv, attn_mask, causal=causal, dropout_p=drop_rate, softmax_scale=None)
|
||||
x = flash_attn_no_pad(
|
||||
qkv, attn_mask, causal=causal, dropout_p=drop_rate, softmax_scale=None
|
||||
)
|
||||
|
||||
b, s, a, d = x.shape
|
||||
out = x.reshape(b, s, -1)
|
||||
return out
|
||||
|
||||
|
||||
def tile(x, sp_size):
|
||||
x = rearrange(x, "b (sp t h w) head d -> b (t sp h w) head d", sp=sp_size, t=30 // sp_size, h=48, w=80)
|
||||
return rearrange(x,
|
||||
"b (n_t ts_t n_h ts_h n_w ts_w) h d -> b (n_t n_h n_w ts_t ts_h ts_w) h d",
|
||||
n_t=5,
|
||||
n_h=6,
|
||||
n_w=10,
|
||||
ts_t=6,
|
||||
ts_h=8,
|
||||
ts_w=8)
|
||||
|
||||
|
||||
def untile(x, sp_size):
|
||||
x = rearrange(x,
|
||||
"b (n_t n_h n_w ts_t ts_h ts_w) h d -> b (n_t ts_t n_h ts_h n_w ts_w) h d",
|
||||
n_t=5,
|
||||
n_h=6,
|
||||
n_w=10,
|
||||
ts_t=6,
|
||||
ts_h=8,
|
||||
ts_w=8)
|
||||
return rearrange(x, "b (t sp h w) head d -> b (sp t h w) head d", sp=sp_size, t=30 // sp_size, h=48, w=80)
|
||||
|
||||
|
||||
def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=None):
|
||||
def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask):
|
||||
# 1GPU torch.Size([1, 11264, 24, 128]) tensor([ 0, 11275, 11520], device='cuda:0', dtype=torch.int32)
|
||||
# 2GPU torch.Size([1, 5632, 24, 128]) tensor([ 0, 5643, 5888], device='cuda:0', dtype=torch.int32)
|
||||
query, encoder_query = q
|
||||
key, encoder_key = k
|
||||
value, encoder_value = v
|
||||
text_length = text_mask.sum()
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
# batch_size, seq_len, attn_heads, head_dim
|
||||
query = all_to_all_4D(query, scatter_dim=2, gather_dim=1)
|
||||
@@ -72,7 +43,9 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=
|
||||
|
||||
def shrink_head(encoder_state, dim):
|
||||
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
|
||||
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
|
||||
return encoder_state.narrow(
|
||||
dim, nccl_info.rank_within_group * local_heads, local_heads
|
||||
)
|
||||
|
||||
encoder_query = shrink_head(encoder_query, dim=2)
|
||||
encoder_key = shrink_head(encoder_key, dim=2)
|
||||
@@ -82,37 +55,24 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask, mask_strategy=
|
||||
sequence_length = query.size(1)
|
||||
encoder_sequence_length = encoder_query.size(1)
|
||||
|
||||
if mask_strategy[0] is not None:
|
||||
query = torch.cat([tile(query, nccl_info.sp_size), encoder_query], dim=1).transpose(1, 2)
|
||||
key = torch.cat([tile(key, nccl_info.sp_size), encoder_key], dim=1).transpose(1, 2)
|
||||
value = torch.cat([tile(value, nccl_info.sp_size), encoder_value], dim=1).transpose(1, 2)
|
||||
# Hint: please check encoder_query.shape
|
||||
query = torch.cat([query, encoder_query], dim=1)
|
||||
key = torch.cat([key, encoder_key], dim=1)
|
||||
value = torch.cat([value, encoder_value], dim=1)
|
||||
# B, S, 3, H, D
|
||||
qkv = torch.stack([query, key, value], dim=2)
|
||||
|
||||
head_num = query.size(1)
|
||||
current_rank = nccl_info.rank_within_group
|
||||
start_head = current_rank * head_num
|
||||
windows = [mask_strategy[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)
|
||||
else:
|
||||
query = torch.cat([query, encoder_query], dim=1)
|
||||
key = torch.cat([key, encoder_key], dim=1)
|
||||
value = torch.cat([value, encoder_value], dim=1)
|
||||
# B, S, 3, H, D
|
||||
qkv = torch.stack([query, key, value], dim=2)
|
||||
|
||||
attn_mask = F.pad(text_mask, (sequence_length, 0), value=True)
|
||||
hidden_states = flash_attn_no_pad(qkv, attn_mask, causal=False, dropout_p=0.0, softmax_scale=None)
|
||||
|
||||
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes((sequence_length, encoder_sequence_length),
|
||||
dim=1)
|
||||
|
||||
if mask_strategy[0] is not None:
|
||||
hidden_states = untile(hidden_states, nccl_info.sp_size)
|
||||
attn_mask = F.pad(text_mask, (sequence_length, 0), value=True)
|
||||
hidden_states = flash_attn_no_pad(
|
||||
qkv, attn_mask, causal=False, dropout_p=0.0, softmax_scale=None
|
||||
)
|
||||
|
||||
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
|
||||
(sequence_length, encoder_sequence_length), dim=1
|
||||
)
|
||||
if get_sequence_parallel_state():
|
||||
hidden_states = all_to_all_4D(hidden_states, scatter_dim=1, gather_dim=2)
|
||||
encoder_hidden_states = all_gather(encoder_hidden_states, dim=2).contiguous()
|
||||
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
encoder_hidden_states = encoder_hidden_states.to(query.dtype)
|
||||
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange, repeat
|
||||
|
||||
from ..utils.helpers import to_2tuple
|
||||
|
||||
@@ -105,8 +105,11 @@ def timestep_embedding(t, dim, max_period=10000):
|
||||
.. ref_link: https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
|
||||
"""
|
||||
half = dim // 2
|
||||
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) /
|
||||
half).to(device=t.device)
|
||||
freqs = torch.exp(
|
||||
-math.log(max_period)
|
||||
* torch.arange(start=0, end=half, dtype=torch.float32)
|
||||
/ half
|
||||
).to(device=t.device)
|
||||
args = t[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
@@ -137,7 +140,9 @@ class TimestepEmbedder(nn.Module):
|
||||
out_size = hidden_size
|
||||
|
||||
self.mlp = nn.Sequential(
|
||||
nn.Linear(frequency_embedding_size, hidden_size, bias=True, **factory_kwargs),
|
||||
nn.Linear(
|
||||
frequency_embedding_size, hidden_size, bias=True, **factory_kwargs
|
||||
),
|
||||
act_layer(),
|
||||
nn.Linear(hidden_size, out_size, bias=True, **factory_kwargs),
|
||||
)
|
||||
@@ -145,6 +150,8 @@ class TimestepEmbedder(nn.Module):
|
||||
nn.init.normal_(self.mlp[2].weight, std=0.02)
|
||||
|
||||
def forward(self, t):
|
||||
t_freq = timestep_embedding(t, self.frequency_embedding_size, self.max_period).type(self.mlp[0].weight.dtype)
|
||||
t_freq = timestep_embedding(
|
||||
t, self.frequency_embedding_size, self.max_period
|
||||
).type(self.mlp[0].weight.dtype)
|
||||
t_emb = self.mlp(t_freq)
|
||||
return t_emb
|
||||
|
||||
@@ -6,8 +6,8 @@ from functools import partial
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from ..utils.helpers import to_2tuple
|
||||
from .modulate_layers import modulate
|
||||
from ..utils.helpers import to_2tuple
|
||||
|
||||
|
||||
class MLP(nn.Module):
|
||||
@@ -34,11 +34,19 @@ class MLP(nn.Module):
|
||||
drop_probs = to_2tuple(drop)
|
||||
linear_layer = partial(nn.Conv2d, kernel_size=1) if use_conv else nn.Linear
|
||||
|
||||
self.fc1 = linear_layer(in_channels, hidden_channels, bias=bias[0], **factory_kwargs)
|
||||
self.fc1 = linear_layer(
|
||||
in_channels, hidden_channels, bias=bias[0], **factory_kwargs
|
||||
)
|
||||
self.act = act_layer()
|
||||
self.drop1 = nn.Dropout(drop_probs[0])
|
||||
self.norm = (norm_layer(hidden_channels, **factory_kwargs) if norm_layer is not None else nn.Identity())
|
||||
self.fc2 = linear_layer(hidden_channels, out_features, bias=bias[1], **factory_kwargs)
|
||||
self.norm = (
|
||||
norm_layer(hidden_channels, **factory_kwargs)
|
||||
if norm_layer is not None
|
||||
else nn.Identity()
|
||||
)
|
||||
self.fc2 = linear_layer(
|
||||
hidden_channels, out_features, bias=bias[1], **factory_kwargs
|
||||
)
|
||||
self.drop2 = nn.Dropout(drop_probs[1])
|
||||
|
||||
def forward(self, x):
|
||||
@@ -69,12 +77,16 @@ class MLPEmbedder(nn.Module):
|
||||
class FinalLayer(nn.Module):
|
||||
"""The final layer of DiT."""
|
||||
|
||||
def __init__(self, hidden_size, patch_size, out_channels, act_layer, device=None, dtype=None):
|
||||
def __init__(
|
||||
self, hidden_size, patch_size, out_channels, act_layer, device=None, dtype=None
|
||||
):
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
|
||||
# Just use LayerNorm for the final layer
|
||||
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
|
||||
self.norm_final = nn.LayerNorm(
|
||||
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
|
||||
)
|
||||
if isinstance(patch_size, int):
|
||||
self.linear = nn.Linear(
|
||||
hidden_size,
|
||||
|
||||
@@ -1,27 +1,29 @@
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
from typing import Any, List, Tuple, Optional, Union, Dict
|
||||
from einops import rearrange
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.models import ModelMixin
|
||||
from einops import rearrange
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.models.hunyuan.modules.posemb_layers import get_nd_rotary_pos_embed
|
||||
from fastvideo.utils.parallel_states import nccl_info
|
||||
from diffusers.models import ModelMixin
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
|
||||
from .activation_layers import get_activation_layer
|
||||
from .attenion import parallel_attention
|
||||
from .embed_layers import PatchEmbed, TextProjection, TimestepEmbedder
|
||||
from .mlp_layers import MLP, FinalLayer, MLPEmbedder
|
||||
from .modulate_layers import ModulateDiT, apply_gate, modulate
|
||||
from .norm_layers import get_norm_layer
|
||||
from .embed_layers import TimestepEmbedder, PatchEmbed, TextProjection
|
||||
from .attenion import parallel_attention
|
||||
from .posemb_layers import apply_rotary_emb
|
||||
from .mlp_layers import MLP, MLPEmbedder, FinalLayer
|
||||
from .modulate_layers import ModulateDiT, modulate, apply_gate
|
||||
from .token_refiner import SingleTokenRefiner
|
||||
from fastvideo.models.hunyuan.modules.posemb_layers import get_nd_rotary_pos_embed
|
||||
|
||||
from fastvideo.utils.parallel_states import nccl_info
|
||||
|
||||
|
||||
class MMDoubleStreamBlock(nn.Module):
|
||||
"""
|
||||
A multimodal dit block with separate modulation for
|
||||
A multimodal dit block with seperate modulation for
|
||||
text and image/video, see more details (SD3): https://arxiv.org/abs/2403.03206
|
||||
(Flux.1): https://github.com/black-forest-labs/flux
|
||||
"""
|
||||
@@ -52,17 +54,31 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
act_layer=get_activation_layer("silu"),
|
||||
**factory_kwargs,
|
||||
)
|
||||
self.img_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
|
||||
self.img_norm1 = nn.LayerNorm(
|
||||
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
|
||||
)
|
||||
|
||||
self.img_attn_qkv = nn.Linear(hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs)
|
||||
self.img_attn_qkv = nn.Linear(
|
||||
hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs
|
||||
)
|
||||
qk_norm_layer = get_norm_layer(qk_norm_type)
|
||||
self.img_attn_q_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm else nn.Identity())
|
||||
self.img_attn_k_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm else nn.Identity())
|
||||
self.img_attn_proj = nn.Linear(hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs)
|
||||
self.img_attn_q_norm = (
|
||||
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm
|
||||
else nn.Identity()
|
||||
)
|
||||
self.img_attn_k_norm = (
|
||||
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm
|
||||
else nn.Identity()
|
||||
)
|
||||
self.img_attn_proj = nn.Linear(
|
||||
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
|
||||
)
|
||||
|
||||
self.img_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
|
||||
self.img_norm2 = nn.LayerNorm(
|
||||
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
|
||||
)
|
||||
self.img_mlp = MLP(
|
||||
hidden_size,
|
||||
mlp_hidden_dim,
|
||||
@@ -77,16 +93,30 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
act_layer=get_activation_layer("silu"),
|
||||
**factory_kwargs,
|
||||
)
|
||||
self.txt_norm1 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
|
||||
self.txt_norm1 = nn.LayerNorm(
|
||||
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
|
||||
)
|
||||
|
||||
self.txt_attn_qkv = nn.Linear(hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs)
|
||||
self.txt_attn_q_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm else nn.Identity())
|
||||
self.txt_attn_k_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm else nn.Identity())
|
||||
self.txt_attn_proj = nn.Linear(hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs)
|
||||
self.txt_attn_qkv = nn.Linear(
|
||||
hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs
|
||||
)
|
||||
self.txt_attn_q_norm = (
|
||||
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm
|
||||
else nn.Identity()
|
||||
)
|
||||
self.txt_attn_k_norm = (
|
||||
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm
|
||||
else nn.Identity()
|
||||
)
|
||||
self.txt_attn_proj = nn.Linear(
|
||||
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
|
||||
)
|
||||
|
||||
self.txt_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
|
||||
self.txt_norm2 = nn.LayerNorm(
|
||||
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
|
||||
)
|
||||
self.txt_mlp = MLP(
|
||||
hidden_size,
|
||||
mlp_hidden_dim,
|
||||
@@ -109,7 +139,6 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
vec: torch.Tensor,
|
||||
freqs_cis: tuple = None,
|
||||
text_mask: torch.Tensor = None,
|
||||
mask_strategy=None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
(
|
||||
img_mod1_shift,
|
||||
@@ -130,9 +159,13 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
|
||||
# Prepare image for attention.
|
||||
img_modulated = self.img_norm1(img)
|
||||
img_modulated = modulate(img_modulated, shift=img_mod1_shift, scale=img_mod1_scale)
|
||||
img_modulated = modulate(
|
||||
img_modulated, shift=img_mod1_shift, scale=img_mod1_scale
|
||||
)
|
||||
img_qkv = self.img_attn_qkv(img_modulated)
|
||||
img_q, img_k, img_v = rearrange(img_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
|
||||
img_q, img_k, img_v = rearrange(
|
||||
img_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
|
||||
)
|
||||
# Apply QK-Norm if needed
|
||||
img_q = self.img_attn_q_norm(img_q).to(img_v)
|
||||
img_k = self.img_attn_k_norm(img_k).to(img_v)
|
||||
@@ -142,7 +175,9 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
|
||||
def shrink_head(encoder_state, dim):
|
||||
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
|
||||
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
|
||||
return encoder_state.narrow(
|
||||
dim, nccl_info.rank_within_group * local_heads, local_heads
|
||||
)
|
||||
|
||||
freqs_cis = (
|
||||
shrink_head(freqs_cis[0], dim=0),
|
||||
@@ -150,15 +185,20 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
)
|
||||
|
||||
img_qq, img_kk = apply_rotary_emb(img_q, img_k, freqs_cis, head_first=False)
|
||||
assert (img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
|
||||
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
|
||||
assert (
|
||||
img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
|
||||
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
|
||||
img_q, img_k = img_qq, img_kk
|
||||
|
||||
# Prepare txt for attention.
|
||||
txt_modulated = self.txt_norm1(txt)
|
||||
txt_modulated = modulate(txt_modulated, shift=txt_mod1_shift, scale=txt_mod1_scale)
|
||||
txt_modulated = modulate(
|
||||
txt_modulated, shift=txt_mod1_shift, scale=txt_mod1_scale
|
||||
)
|
||||
txt_qkv = self.txt_attn_qkv(txt_modulated)
|
||||
txt_q, txt_k, txt_v = rearrange(txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
|
||||
txt_q, txt_k, txt_v = rearrange(
|
||||
txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
|
||||
)
|
||||
# Apply QK-Norm if needed.
|
||||
txt_q = self.txt_attn_q_norm(txt_q).to(txt_v)
|
||||
txt_k = self.txt_attn_k_norm(txt_k).to(txt_v)
|
||||
@@ -170,26 +210,34 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
img_q_len=img_q.shape[1],
|
||||
img_kv_len=img_k.shape[1],
|
||||
text_mask=text_mask,
|
||||
mask_strategy=mask_strategy,
|
||||
)
|
||||
|
||||
# attention computation end
|
||||
|
||||
img_attn, txt_attn = attn[:, :img.shape[1]], attn[:, img.shape[1]:]
|
||||
img_attn, txt_attn = attn[:, : img.shape[1]], attn[:, img.shape[1] :]
|
||||
|
||||
# Calculate the img blocks.
|
||||
# Calculate the img bloks.
|
||||
img = img + apply_gate(self.img_attn_proj(img_attn), gate=img_mod1_gate)
|
||||
img = img + apply_gate(
|
||||
self.img_mlp(modulate(self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale)),
|
||||
self.img_mlp(
|
||||
modulate(
|
||||
self.img_norm2(img), shift=img_mod2_shift, scale=img_mod2_scale
|
||||
)
|
||||
),
|
||||
gate=img_mod2_gate,
|
||||
)
|
||||
|
||||
# Calculate the txt blocks.
|
||||
# Calculate the txt bloks.
|
||||
txt = txt + apply_gate(self.txt_attn_proj(txt_attn), gate=txt_mod1_gate)
|
||||
txt = txt + apply_gate(
|
||||
self.txt_mlp(modulate(self.txt_norm2(txt), shift=txt_mod2_shift, scale=txt_mod2_scale)),
|
||||
self.txt_mlp(
|
||||
modulate(
|
||||
self.txt_norm2(txt), shift=txt_mod2_shift, scale=txt_mod2_scale
|
||||
)
|
||||
),
|
||||
gate=txt_mod2_gate,
|
||||
)
|
||||
|
||||
return img, txt
|
||||
|
||||
|
||||
@@ -222,20 +270,32 @@ class MMSingleStreamBlock(nn.Module):
|
||||
head_dim = hidden_size // heads_num
|
||||
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
|
||||
self.mlp_hidden_dim = mlp_hidden_dim
|
||||
self.scale = qk_scale or head_dim**-0.5
|
||||
self.scale = qk_scale or head_dim ** -0.5
|
||||
|
||||
# qkv and mlp_in
|
||||
self.linear1 = nn.Linear(hidden_size, hidden_size * 3 + mlp_hidden_dim, **factory_kwargs)
|
||||
self.linear1 = nn.Linear(
|
||||
hidden_size, hidden_size * 3 + mlp_hidden_dim, **factory_kwargs
|
||||
)
|
||||
# proj and mlp_out
|
||||
self.linear2 = nn.Linear(hidden_size + mlp_hidden_dim, hidden_size, **factory_kwargs)
|
||||
self.linear2 = nn.Linear(
|
||||
hidden_size + mlp_hidden_dim, hidden_size, **factory_kwargs
|
||||
)
|
||||
|
||||
qk_norm_layer = get_norm_layer(qk_norm_type)
|
||||
self.q_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm else nn.Identity())
|
||||
self.k_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm else nn.Identity())
|
||||
self.q_norm = (
|
||||
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm
|
||||
else nn.Identity()
|
||||
)
|
||||
self.k_norm = (
|
||||
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm
|
||||
else nn.Identity()
|
||||
)
|
||||
|
||||
self.pre_norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
|
||||
self.pre_norm = nn.LayerNorm(
|
||||
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
|
||||
)
|
||||
|
||||
self.mlp_act = get_activation_layer(mlp_act_type)()
|
||||
self.modulation = ModulateDiT(
|
||||
@@ -259,11 +319,12 @@ class MMSingleStreamBlock(nn.Module):
|
||||
txt_len: int,
|
||||
freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None,
|
||||
text_mask: torch.Tensor = None,
|
||||
mask_strategy=None,
|
||||
) -> torch.Tensor:
|
||||
mod_shift, mod_scale, mod_gate = self.modulation(vec).chunk(3, dim=-1)
|
||||
x_mod = modulate(self.pre_norm(x), shift=mod_shift, scale=mod_scale)
|
||||
qkv, mlp = torch.split(self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1)
|
||||
qkv, mlp = torch.split(
|
||||
self.linear1(x_mod), [3 * self.hidden_size, self.mlp_hidden_dim], dim=-1
|
||||
)
|
||||
|
||||
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
|
||||
|
||||
@@ -273,19 +334,19 @@ class MMSingleStreamBlock(nn.Module):
|
||||
|
||||
def shrink_head(encoder_state, dim):
|
||||
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
|
||||
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
|
||||
return encoder_state.narrow(
|
||||
dim, nccl_info.rank_within_group * local_heads, local_heads
|
||||
)
|
||||
|
||||
freqs_cis = (
|
||||
shrink_head(freqs_cis[0], dim=0),
|
||||
shrink_head(freqs_cis[1], dim=0),
|
||||
)
|
||||
freqs_cis = (shrink_head(freqs_cis[0], dim=0), shrink_head(freqs_cis[1], dim=0))
|
||||
|
||||
img_q, txt_q = q[:, :-txt_len, :, :], q[:, -txt_len:, :, :]
|
||||
img_k, txt_k = k[:, :-txt_len, :, :], k[:, -txt_len:, :, :]
|
||||
img_v, txt_v = v[:, :-txt_len, :, :], v[:, -txt_len:, :, :]
|
||||
img_qq, img_kk = apply_rotary_emb(img_q, img_k, freqs_cis, head_first=False)
|
||||
assert (img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
|
||||
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
|
||||
assert (
|
||||
img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
|
||||
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
|
||||
img_q, img_k = img_qq, img_kk
|
||||
|
||||
attn = parallel_attention(
|
||||
@@ -295,7 +356,6 @@ class MMSingleStreamBlock(nn.Module):
|
||||
img_q_len=img_q.shape[1],
|
||||
img_kv_len=img_k.shape[1],
|
||||
text_mask=text_mask,
|
||||
mask_strategy=mask_strategy,
|
||||
)
|
||||
|
||||
# attention computation end
|
||||
@@ -398,15 +458,21 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
self.text_projection = text_projection
|
||||
|
||||
if hidden_size % heads_num != 0:
|
||||
raise ValueError(f"Hidden size {hidden_size} must be divisible by heads_num {heads_num}")
|
||||
raise ValueError(
|
||||
f"Hidden size {hidden_size} must be divisible by heads_num {heads_num}"
|
||||
)
|
||||
pe_dim = hidden_size // heads_num
|
||||
if sum(rope_dim_list) != pe_dim:
|
||||
raise ValueError(f"Got {rope_dim_list} but expected positional dim {pe_dim}")
|
||||
raise ValueError(
|
||||
f"Got {rope_dim_list} but expected positional dim {pe_dim}"
|
||||
)
|
||||
self.hidden_size = hidden_size
|
||||
self.heads_num = heads_num
|
||||
|
||||
# image projection
|
||||
self.img_in = PatchEmbed(self.patch_size, self.in_channels, self.hidden_size, **factory_kwargs)
|
||||
self.img_in = PatchEmbed(
|
||||
self.patch_size, self.in_channels, self.hidden_size, **factory_kwargs
|
||||
)
|
||||
|
||||
# text projection
|
||||
if self.text_projection == "linear":
|
||||
@@ -425,44 +491,61 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
**factory_kwargs,
|
||||
)
|
||||
else:
|
||||
raise NotImplementedError(f"Unsupported text_projection: {self.text_projection}")
|
||||
raise NotImplementedError(
|
||||
f"Unsupported text_projection: {self.text_projection}"
|
||||
)
|
||||
|
||||
# time modulation
|
||||
self.time_in = TimestepEmbedder(self.hidden_size, get_activation_layer("silu"), **factory_kwargs)
|
||||
self.time_in = TimestepEmbedder(
|
||||
self.hidden_size, get_activation_layer("silu"), **factory_kwargs
|
||||
)
|
||||
|
||||
# text modulation
|
||||
self.vector_in = MLPEmbedder(self.config.text_states_dim_2, self.hidden_size, **factory_kwargs)
|
||||
self.vector_in = MLPEmbedder(
|
||||
self.config.text_states_dim_2, self.hidden_size, **factory_kwargs
|
||||
)
|
||||
|
||||
# guidance modulation
|
||||
self.guidance_in = (TimestepEmbedder(self.hidden_size, get_activation_layer("silu"), **factory_kwargs)
|
||||
if guidance_embed else None)
|
||||
self.guidance_in = (
|
||||
TimestepEmbedder(
|
||||
self.hidden_size, get_activation_layer("silu"), **factory_kwargs
|
||||
)
|
||||
if guidance_embed
|
||||
else None
|
||||
)
|
||||
|
||||
# double blocks
|
||||
self.double_blocks = nn.ModuleList([
|
||||
MMDoubleStreamBlock(
|
||||
self.hidden_size,
|
||||
self.heads_num,
|
||||
mlp_width_ratio=mlp_width_ratio,
|
||||
mlp_act_type=mlp_act_type,
|
||||
qk_norm=qk_norm,
|
||||
qk_norm_type=qk_norm_type,
|
||||
qkv_bias=qkv_bias,
|
||||
**factory_kwargs,
|
||||
) for _ in range(mm_double_blocks_depth)
|
||||
])
|
||||
self.double_blocks = nn.ModuleList(
|
||||
[
|
||||
MMDoubleStreamBlock(
|
||||
self.hidden_size,
|
||||
self.heads_num,
|
||||
mlp_width_ratio=mlp_width_ratio,
|
||||
mlp_act_type=mlp_act_type,
|
||||
qk_norm=qk_norm,
|
||||
qk_norm_type=qk_norm_type,
|
||||
qkv_bias=qkv_bias,
|
||||
**factory_kwargs,
|
||||
)
|
||||
for _ in range(mm_double_blocks_depth)
|
||||
]
|
||||
)
|
||||
|
||||
# single blocks
|
||||
self.single_blocks = nn.ModuleList([
|
||||
MMSingleStreamBlock(
|
||||
self.hidden_size,
|
||||
self.heads_num,
|
||||
mlp_width_ratio=mlp_width_ratio,
|
||||
mlp_act_type=mlp_act_type,
|
||||
qk_norm=qk_norm,
|
||||
qk_norm_type=qk_norm_type,
|
||||
**factory_kwargs,
|
||||
) for _ in range(mm_single_blocks_depth)
|
||||
])
|
||||
self.single_blocks = nn.ModuleList(
|
||||
[
|
||||
MMSingleStreamBlock(
|
||||
self.hidden_size,
|
||||
self.heads_num,
|
||||
mlp_width_ratio=mlp_width_ratio,
|
||||
mlp_act_type=mlp_act_type,
|
||||
qk_norm=qk_norm,
|
||||
qk_norm_type=qk_norm_type,
|
||||
**factory_kwargs,
|
||||
)
|
||||
for _ in range(mm_single_blocks_depth)
|
||||
]
|
||||
)
|
||||
|
||||
self.final_layer = FinalLayer(
|
||||
self.hidden_size,
|
||||
@@ -486,12 +569,14 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
|
||||
def get_rotary_pos_embed(self, rope_sizes):
|
||||
target_ndim = 3
|
||||
|
||||
ndim = 5 - 2
|
||||
head_dim = self.hidden_size // self.heads_num
|
||||
rope_dim_list = self.rope_dim_list
|
||||
if rope_dim_list is None:
|
||||
rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)]
|
||||
assert (sum(rope_dim_list) == head_dim), "sum(rope_dim_list) should equal to head_dim of attention layer"
|
||||
assert (
|
||||
sum(rope_dim_list) == head_dim
|
||||
), "sum(rope_dim_list) should equal to head_dim of attention layer"
|
||||
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
|
||||
rope_dim_list,
|
||||
rope_sizes,
|
||||
@@ -514,27 +599,27 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
mask_strategy=None,
|
||||
output_features=False,
|
||||
output_features_stride=8,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
return_dict: bool = False,
|
||||
guidance=None,
|
||||
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||
if guidance is None:
|
||||
guidance = torch.tensor([6016.0], device=hidden_states.device, dtype=torch.bfloat16)
|
||||
if mask_strategy is None:
|
||||
mask_strategy = [[None] * self.heads_num for _ in range(len(self.double_blocks) + len(self.single_blocks))]
|
||||
if guidance == None:
|
||||
guidance = torch.tensor(
|
||||
[6016.0], device=hidden_states.device, dtype=torch.bfloat16
|
||||
)
|
||||
out = {}
|
||||
img = x = hidden_states
|
||||
text_mask = encoder_attention_mask
|
||||
t = timestep
|
||||
txt = encoder_hidden_states[:, 1:]
|
||||
text_states_2 = encoder_hidden_states[:, 0, :self.config.text_states_dim_2]
|
||||
_, _, ot, oh, ow = x.shape # codespell:ignore
|
||||
text_states_2 = encoder_hidden_states[:, 0, : self.config.text_states_dim_2]
|
||||
_, _, ot, oh, ow = x.shape
|
||||
tt, th, tw = (
|
||||
ot // self.patch_size[0], # codespell:ignore
|
||||
oh // self.patch_size[1], # codespell:ignore
|
||||
ow // self.patch_size[2], # codespell:ignore
|
||||
ot // self.patch_size[0],
|
||||
oh // self.patch_size[1],
|
||||
ow // self.patch_size[2],
|
||||
)
|
||||
original_tt = nccl_info.sp_size * tt
|
||||
freqs_cos, freqs_sin = self.get_rotary_pos_embed((original_tt, th, tw))
|
||||
@@ -547,7 +632,9 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
# guidance modulation
|
||||
if self.guidance_embed:
|
||||
if guidance is None:
|
||||
raise ValueError("Didn't get guidance strength for guidance distilled model.")
|
||||
raise ValueError(
|
||||
"Didn't get guidance strength for guidance distilled model."
|
||||
)
|
||||
|
||||
# our timestep_embedding is merged into guidance_in(TimestepEmbedder)
|
||||
vec = vec + self.guidance_in(guidance)
|
||||
@@ -559,31 +646,34 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
elif self.text_projection == "single_refiner":
|
||||
txt = self.txt_in(txt, t, text_mask if self.use_attention_mask else None)
|
||||
else:
|
||||
raise NotImplementedError(f"Unsupported text_projection: {self.text_projection}")
|
||||
raise NotImplementedError(
|
||||
f"Unsupported text_projection: {self.text_projection}"
|
||||
)
|
||||
|
||||
txt_seq_len = txt.shape[1]
|
||||
img_seq_len = img.shape[1]
|
||||
|
||||
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
|
||||
# --------------------- Pass through DiT blocks ------------------------
|
||||
for _, block in enumerate(self.double_blocks):
|
||||
double_block_args = [img, txt, vec, freqs_cis, text_mask]
|
||||
|
||||
for index, block in enumerate(self.double_blocks):
|
||||
double_block_args = [img, txt, vec, freqs_cis, text_mask, mask_strategy[index]]
|
||||
img, txt = block(*double_block_args)
|
||||
|
||||
# Merge txt and img to pass through single stream blocks.
|
||||
x = torch.cat((img, txt), 1)
|
||||
if output_features:
|
||||
features_list = []
|
||||
if len(self.single_blocks) > 0:
|
||||
for index, block in enumerate(self.single_blocks):
|
||||
for _, block in enumerate(self.single_blocks):
|
||||
single_block_args = [
|
||||
x,
|
||||
vec,
|
||||
txt_seq_len,
|
||||
(freqs_cos, freqs_sin),
|
||||
text_mask,
|
||||
mask_strategy[index + len(self.double_blocks)],
|
||||
]
|
||||
|
||||
x = block(*single_block_args)
|
||||
if output_features and _ % output_features_stride == 0:
|
||||
features_list.append(x[:, :img_seq_len, ...])
|
||||
@@ -594,7 +684,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
|
||||
|
||||
img = self.unpatchify(img, tt, th, tw)
|
||||
assert not return_dict, "return_dict is not supported."
|
||||
assert return_dict == False, "return_dict is not supported."
|
||||
if output_features:
|
||||
features_list = torch.stack(features_list, dim=0)
|
||||
else:
|
||||
@@ -618,24 +708,25 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
|
||||
def params_count(self):
|
||||
counts = {
|
||||
"double":
|
||||
sum([
|
||||
sum(p.numel()
|
||||
for p in block.img_attn_qkv.parameters()) + sum(p.numel()
|
||||
for p in block.img_attn_proj.parameters()) +
|
||||
sum(p.numel() for p in block.img_mlp.parameters()) + sum(p.numel()
|
||||
for p in block.txt_attn_qkv.parameters()) +
|
||||
sum(p.numel() for p in block.txt_attn_proj.parameters()) + sum(p.numel()
|
||||
for p in block.txt_mlp.parameters())
|
||||
for block in self.double_blocks
|
||||
]),
|
||||
"single":
|
||||
sum([
|
||||
sum(p.numel() for p in block.linear1.parameters()) + sum(p.numel() for p in block.linear2.parameters())
|
||||
for block in self.single_blocks
|
||||
]),
|
||||
"total":
|
||||
sum(p.numel() for p in self.parameters()),
|
||||
"double": sum(
|
||||
[
|
||||
sum(p.numel() for p in block.img_attn_qkv.parameters())
|
||||
+ sum(p.numel() for p in block.img_attn_proj.parameters())
|
||||
+ sum(p.numel() for p in block.img_mlp.parameters())
|
||||
+ sum(p.numel() for p in block.txt_attn_qkv.parameters())
|
||||
+ sum(p.numel() for p in block.txt_attn_proj.parameters())
|
||||
+ sum(p.numel() for p in block.txt_mlp.parameters())
|
||||
for block in self.double_blocks
|
||||
]
|
||||
),
|
||||
"single": sum(
|
||||
[
|
||||
sum(p.numel() for p in block.linear1.parameters())
|
||||
+ sum(p.numel() for p in block.linear2.parameters())
|
||||
for block in self.single_blocks
|
||||
]
|
||||
),
|
||||
"total": sum(p.numel() for p in self.parameters()),
|
||||
}
|
||||
counts["attn+mlp"] = counts["double"] + counts["single"]
|
||||
return counts
|
||||
|
||||
@@ -18,7 +18,9 @@ class ModulateDiT(nn.Module):
|
||||
factory_kwargs = {"dtype": dtype, "device": device}
|
||||
super().__init__()
|
||||
self.act = act_layer()
|
||||
self.linear = nn.Linear(hidden_size, factor * hidden_size, bias=True, **factory_kwargs)
|
||||
self.linear = nn.Linear(
|
||||
hidden_size, factor * hidden_size, bias=True, **factory_kwargs
|
||||
)
|
||||
# Zero-initialize the modulation
|
||||
nn.init.zeros_(self.linear.weight)
|
||||
nn.init.zeros_(self.linear.bias)
|
||||
@@ -68,7 +70,6 @@ def apply_gate(x, gate=None, tanh=False):
|
||||
|
||||
|
||||
def ckpt_wrapper(module):
|
||||
|
||||
def ckpt_forward(*inputs):
|
||||
outputs = module(*inputs)
|
||||
return outputs
|
||||
@@ -76,8 +77,11 @@ def ckpt_wrapper(module):
|
||||
return ckpt_forward
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
|
||||
@@ -3,7 +3,6 @@ import torch.nn as nn
|
||||
|
||||
|
||||
class RMSNorm(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
|
||||
@@ -1,11 +1,10 @@
|
||||
from typing import List, Tuple, Union
|
||||
|
||||
import torch
|
||||
from typing import Union, Tuple, List
|
||||
|
||||
|
||||
def _to_tuple(x, dim=2):
|
||||
if isinstance(x, int):
|
||||
return (x, ) * dim
|
||||
return (x,) * dim
|
||||
elif len(x) == dim:
|
||||
return x
|
||||
else:
|
||||
@@ -30,7 +29,7 @@ def get_meshgrid_nd(start, *args, dim=2):
|
||||
if len(args) == 0:
|
||||
# start is grid_size
|
||||
num = _to_tuple(start, dim=dim)
|
||||
start = (0, ) * dim
|
||||
start = (0,) * dim
|
||||
stop = num
|
||||
elif len(args) == 1:
|
||||
# start is start, args[0] is stop, step is 1
|
||||
@@ -100,7 +99,10 @@ def reshape_for_broadcast(
|
||||
x.shape[-2],
|
||||
x.shape[-1],
|
||||
), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}"
|
||||
shape = [d if i == ndim - 2 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
|
||||
shape = [
|
||||
d if i == ndim - 2 or i == ndim - 1 else 1
|
||||
for i, d in enumerate(x.shape)
|
||||
]
|
||||
else:
|
||||
assert freqs_cis[0].shape == (
|
||||
x.shape[1],
|
||||
@@ -115,7 +117,10 @@ def reshape_for_broadcast(
|
||||
x.shape[-2],
|
||||
x.shape[-1],
|
||||
), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}"
|
||||
shape = [d if i == ndim - 2 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
|
||||
shape = [
|
||||
d if i == ndim - 2 or i == ndim - 1 else 1
|
||||
for i, d in enumerate(x.shape)
|
||||
]
|
||||
else:
|
||||
assert freqs_cis.shape == (
|
||||
x.shape[1],
|
||||
@@ -126,7 +131,9 @@ def reshape_for_broadcast(
|
||||
|
||||
|
||||
def rotate_half(x):
|
||||
x_real, x_imag = (x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1)) # [B, S, H, D//2]
|
||||
x_real, x_imag = (
|
||||
x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1)
|
||||
) # [B, S, H, D//2]
|
||||
return torch.stack([-x_imag, x_real], dim=-1).flatten(3)
|
||||
|
||||
|
||||
@@ -164,12 +171,18 @@ def apply_rotary_emb(
|
||||
xk_out = (xk.float() * cos + rotate_half(xk.float()) * sin).type_as(xk)
|
||||
else:
|
||||
# view_as_complex will pack [..., D/2, 2](real) to [..., D/2](complex)
|
||||
xq_ = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2)) # [B, S, H, D//2]
|
||||
freqs_cis = reshape_for_broadcast(freqs_cis, xq_, head_first).to(xq.device) # [S, D//2] --> [1, S, 1, D//2]
|
||||
xq_ = torch.view_as_complex(
|
||||
xq.float().reshape(*xq.shape[:-1], -1, 2)
|
||||
) # [B, S, H, D//2]
|
||||
freqs_cis = reshape_for_broadcast(freqs_cis, xq_, head_first).to(
|
||||
xq.device
|
||||
) # [S, D//2] --> [1, S, 1, D//2]
|
||||
# (real, imag) * (cos, sin) = (real * cos - imag * sin, imag * cos + real * sin)
|
||||
# view_as_real will expand [..., D/2](complex) to [..., D/2, 2](real)
|
||||
xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3).type_as(xq)
|
||||
xk_ = torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2)) # [B, S, H, D//2]
|
||||
xk_ = torch.view_as_complex(
|
||||
xk.float().reshape(*xk.shape[:-1], -1, 2)
|
||||
) # [B, S, H, D//2]
|
||||
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3).type_as(xk)
|
||||
|
||||
return xq_out, xk_out
|
||||
@@ -203,21 +216,25 @@ def get_nd_rotary_pos_embed(
|
||||
pos_embed (torch.Tensor): [HW, D/2]
|
||||
"""
|
||||
|
||||
grid = get_meshgrid_nd(start, *args, dim=len(rope_dim_list)) # [3, W, H, D] / [2, W, H]
|
||||
grid = get_meshgrid_nd(
|
||||
start, *args, dim=len(rope_dim_list)
|
||||
) # [3, W, H, D] / [2, W, H]
|
||||
|
||||
if isinstance(theta_rescale_factor, int) or isinstance(theta_rescale_factor, float):
|
||||
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
|
||||
elif isinstance(theta_rescale_factor, list) and len(theta_rescale_factor) == 1:
|
||||
theta_rescale_factor = [theta_rescale_factor[0]] * len(rope_dim_list)
|
||||
assert len(theta_rescale_factor) == len(
|
||||
rope_dim_list), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
|
||||
rope_dim_list
|
||||
), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
|
||||
|
||||
if isinstance(interpolation_factor, int) or isinstance(interpolation_factor, float):
|
||||
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
|
||||
elif isinstance(interpolation_factor, list) and len(interpolation_factor) == 1:
|
||||
interpolation_factor = [interpolation_factor[0]] * len(rope_dim_list)
|
||||
assert len(interpolation_factor) == len(
|
||||
rope_dim_list), "len(interpolation_factor) should equal to len(rope_dim_list)"
|
||||
rope_dim_list
|
||||
), "len(interpolation_factor) should equal to len(rope_dim_list)"
|
||||
|
||||
# use 1/ndim of dimensions to encode grid_axis
|
||||
embs = []
|
||||
@@ -275,9 +292,11 @@ def get_1d_rotary_pos_embed(
|
||||
# proposed by reddit user bloc97, to rescale rotary embeddings to longer sequence length without fine-tuning
|
||||
# has some connection to NTK literature
|
||||
if theta_rescale_factor != 1.0:
|
||||
theta *= theta_rescale_factor**(dim / (dim - 2))
|
||||
theta *= theta_rescale_factor ** (dim / (dim - 2))
|
||||
|
||||
freqs = 1.0 / (theta**(torch.arange(0, dim, 2)[:(dim // 2)].float() / dim)) # [D/2]
|
||||
freqs = 1.0 / (
|
||||
theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)
|
||||
) # [D/2]
|
||||
# assert interpolation_factor == 1.0, f"interpolation_factor: {interpolation_factor}"
|
||||
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
|
||||
if use_real:
|
||||
@@ -285,5 +304,7 @@ def get_1d_rotary_pos_embed(
|
||||
freqs_sin = freqs.sin().repeat_interleave(2, dim=1) # [S, D]
|
||||
return freqs_cos, freqs_sin
|
||||
else:
|
||||
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 # [S, D/2]
|
||||
freqs_cis = torch.polar(
|
||||
torch.ones_like(freqs), freqs
|
||||
) # complex64 # [S, D/2]
|
||||
return freqs_cis
|
||||
|
||||
@@ -1,19 +1,19 @@
|
||||
from typing import Optional
|
||||
|
||||
from einops import rearrange
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
|
||||
from .activation_layers import get_activation_layer
|
||||
from .attenion import attention
|
||||
from .embed_layers import TextProjection, TimestepEmbedder
|
||||
from .mlp_layers import MLP
|
||||
from .modulate_layers import apply_gate
|
||||
from .norm_layers import get_norm_layer
|
||||
from .embed_layers import TimestepEmbedder, TextProjection
|
||||
from .attenion import attention
|
||||
from .mlp_layers import MLP
|
||||
from .modulate_layers import modulate, apply_gate
|
||||
|
||||
|
||||
class IndividualTokenRefinerBlock(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
@@ -33,16 +33,30 @@ class IndividualTokenRefinerBlock(nn.Module):
|
||||
head_dim = hidden_size // heads_num
|
||||
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
|
||||
|
||||
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
self.self_attn_qkv = nn.Linear(hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs)
|
||||
self.norm1 = nn.LayerNorm(
|
||||
hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs
|
||||
)
|
||||
self.self_attn_qkv = nn.Linear(
|
||||
hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs
|
||||
)
|
||||
qk_norm_layer = get_norm_layer(qk_norm_type)
|
||||
self.self_attn_q_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm else nn.Identity())
|
||||
self.self_attn_k_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm else nn.Identity())
|
||||
self.self_attn_proj = nn.Linear(hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs)
|
||||
self.self_attn_q_norm = (
|
||||
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm
|
||||
else nn.Identity()
|
||||
)
|
||||
self.self_attn_k_norm = (
|
||||
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
if qk_norm
|
||||
else nn.Identity()
|
||||
)
|
||||
self.self_attn_proj = nn.Linear(
|
||||
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
|
||||
)
|
||||
|
||||
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs)
|
||||
self.norm2 = nn.LayerNorm(
|
||||
hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs
|
||||
)
|
||||
act_layer = get_activation_layer(act_type)
|
||||
self.mlp = MLP(
|
||||
in_channels=hidden_size,
|
||||
@@ -87,7 +101,6 @@ class IndividualTokenRefinerBlock(nn.Module):
|
||||
|
||||
|
||||
class IndividualTokenRefiner(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
@@ -104,25 +117,25 @@ class IndividualTokenRefiner(nn.Module):
|
||||
):
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
self.blocks = nn.ModuleList([
|
||||
IndividualTokenRefinerBlock(
|
||||
hidden_size=hidden_size,
|
||||
heads_num=heads_num,
|
||||
mlp_width_ratio=mlp_width_ratio,
|
||||
mlp_drop_rate=mlp_drop_rate,
|
||||
act_type=act_type,
|
||||
qk_norm=qk_norm,
|
||||
qk_norm_type=qk_norm_type,
|
||||
qkv_bias=qkv_bias,
|
||||
**factory_kwargs,
|
||||
) for _ in range(depth)
|
||||
])
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
IndividualTokenRefinerBlock(
|
||||
hidden_size=hidden_size,
|
||||
heads_num=heads_num,
|
||||
mlp_width_ratio=mlp_width_ratio,
|
||||
mlp_drop_rate=mlp_drop_rate,
|
||||
act_type=act_type,
|
||||
qk_norm=qk_norm,
|
||||
qk_norm_type=qk_norm_type,
|
||||
qkv_bias=qkv_bias,
|
||||
**factory_kwargs,
|
||||
)
|
||||
for _ in range(depth)
|
||||
]
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
c: torch.LongTensor,
|
||||
mask: Optional[torch.Tensor] = None,
|
||||
self, x: torch.Tensor, c: torch.LongTensor, mask: Optional[torch.Tensor] = None,
|
||||
):
|
||||
mask = mask.clone().bool()
|
||||
# avoid attention weight become NaN
|
||||
@@ -158,13 +171,17 @@ class SingleTokenRefiner(nn.Module):
|
||||
self.attn_mode = attn_mode
|
||||
assert self.attn_mode == "torch", "Only support 'torch' mode for token refiner."
|
||||
|
||||
self.input_embedder = nn.Linear(in_channels, hidden_size, bias=True, **factory_kwargs)
|
||||
self.input_embedder = nn.Linear(
|
||||
in_channels, hidden_size, bias=True, **factory_kwargs
|
||||
)
|
||||
|
||||
act_layer = get_activation_layer(act_type)
|
||||
# Build timestep embedding layer
|
||||
self.t_embedder = TimestepEmbedder(hidden_size, act_layer, **factory_kwargs)
|
||||
# Build context embedding layer
|
||||
self.c_embedder = TextProjection(in_channels, hidden_size, act_layer, **factory_kwargs)
|
||||
self.c_embedder = TextProjection(
|
||||
in_channels, hidden_size, act_layer, **factory_kwargs
|
||||
)
|
||||
|
||||
self.individual_token_refiner = IndividualTokenRefiner(
|
||||
hidden_size=hidden_size,
|
||||
@@ -191,7 +208,9 @@ class SingleTokenRefiner(nn.Module):
|
||||
context_aware_representations = x.mean(dim=1)
|
||||
else:
|
||||
mask_float = mask.float().unsqueeze(-1) # [b, s1, 1]
|
||||
context_aware_representations = (x * mask_float).sum(dim=1) / mask_float.sum(dim=1)
|
||||
context_aware_representations = (x * mask_float).sum(
|
||||
dim=1
|
||||
) / mask_float.sum(dim=1)
|
||||
context_aware_representations = self.c_embedder(context_aware_representations)
|
||||
c = timestep_aware_representations + context_aware_representations
|
||||
|
||||
|
||||
@@ -16,6 +16,7 @@ Given Input:
|
||||
input: "{input}"
|
||||
"""
|
||||
|
||||
|
||||
master_mode_prompt = """Master mode - Video Recaption Task:
|
||||
|
||||
You are a large language model specialized in rewriting video descriptions. Your task is to modify the input description.
|
||||
|
||||
@@ -1,12 +1,14 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Tuple
|
||||
from copy import deepcopy
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from transformers import AutoModel, AutoTokenizer, CLIPTextModel, CLIPTokenizer
|
||||
from transformers import CLIPTextModel, CLIPTokenizer, AutoTokenizer, AutoModel
|
||||
from transformers.utils import ModelOutput
|
||||
|
||||
from ..constants import PRECISION_TO_TYPE, TEXT_ENCODER_PATH, TOKENIZER_PATH
|
||||
from ..constants import TEXT_ENCODER_PATH, TOKENIZER_PATH
|
||||
from ..constants import PRECISION_TO_TYPE
|
||||
|
||||
|
||||
def use_default(value, default):
|
||||
@@ -23,13 +25,17 @@ def load_text_encoder(
|
||||
if text_encoder_path is None:
|
||||
text_encoder_path = TEXT_ENCODER_PATH[text_encoder_type]
|
||||
if logger is not None:
|
||||
logger.info(f"Loading text encoder model ({text_encoder_type}) from: {text_encoder_path}")
|
||||
logger.info(
|
||||
f"Loading text encoder model ({text_encoder_type}) from: {text_encoder_path}"
|
||||
)
|
||||
|
||||
if text_encoder_type == "clipL":
|
||||
text_encoder = CLIPTextModel.from_pretrained(text_encoder_path)
|
||||
text_encoder.final_layer_norm = text_encoder.text_model.final_layer_norm
|
||||
elif text_encoder_type == "llm":
|
||||
text_encoder = AutoModel.from_pretrained(text_encoder_path, low_cpu_mem_usage=True)
|
||||
text_encoder = AutoModel.from_pretrained(
|
||||
text_encoder_path, low_cpu_mem_usage=True
|
||||
)
|
||||
text_encoder.final_layer_norm = text_encoder.norm
|
||||
else:
|
||||
raise ValueError(f"Unsupported text encoder type: {text_encoder_type}")
|
||||
@@ -49,7 +55,9 @@ def load_text_encoder(
|
||||
return text_encoder, text_encoder_path
|
||||
|
||||
|
||||
def load_tokenizer(tokenizer_type, tokenizer_path=None, padding_side="right", logger=None):
|
||||
def load_tokenizer(
|
||||
tokenizer_type, tokenizer_path=None, padding_side="right", logger=None
|
||||
):
|
||||
if tokenizer_path is None:
|
||||
tokenizer_path = TOKENIZER_PATH[tokenizer_type]
|
||||
if logger is not None:
|
||||
@@ -58,7 +66,9 @@ def load_tokenizer(tokenizer_type, tokenizer_path=None, padding_side="right", lo
|
||||
if tokenizer_type == "clipL":
|
||||
tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path, max_length=77)
|
||||
elif tokenizer_type == "llm":
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path, padding_side=padding_side)
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
tokenizer_path, padding_side=padding_side
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported tokenizer type: {tokenizer_type}")
|
||||
|
||||
@@ -90,7 +100,6 @@ class TextEncoderModelOutput(ModelOutput):
|
||||
|
||||
|
||||
class TextEncoder(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
text_encoder_type: str,
|
||||
@@ -115,12 +124,20 @@ class TextEncoder(nn.Module):
|
||||
self.max_length = max_length
|
||||
self.precision = text_encoder_precision
|
||||
self.model_path = text_encoder_path
|
||||
self.tokenizer_type = (tokenizer_type if tokenizer_type is not None else text_encoder_type)
|
||||
self.tokenizer_path = (tokenizer_path if tokenizer_path is not None else text_encoder_path)
|
||||
self.tokenizer_type = (
|
||||
tokenizer_type if tokenizer_type is not None else text_encoder_type
|
||||
)
|
||||
self.tokenizer_path = (
|
||||
tokenizer_path if tokenizer_path is not None else 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.input_max_length = (input_max_length if input_max_length is not None else max_length)
|
||||
assert (
|
||||
use_attention_mask is True
|
||||
), "Attention mask is True required when training videos."
|
||||
self.input_max_length = (
|
||||
input_max_length if input_max_length is not None else max_length
|
||||
)
|
||||
self.prompt_template = prompt_template
|
||||
self.prompt_template_video = prompt_template_video
|
||||
self.hidden_state_skip_layer = hidden_state_skip_layer
|
||||
@@ -130,21 +147,26 @@ class TextEncoder(nn.Module):
|
||||
|
||||
self.use_template = self.prompt_template is not None
|
||||
if self.use_template:
|
||||
assert (isinstance(self.prompt_template, dict) and "template" in self.prompt_template
|
||||
), f"`prompt_template` must be a dictionary with a key 'template', got {self.prompt_template}"
|
||||
assert (
|
||||
isinstance(self.prompt_template, dict)
|
||||
and "template" in self.prompt_template
|
||||
), f"`prompt_template` must be a dictionary with a key 'template', got {self.prompt_template}"
|
||||
assert "{}" in str(self.prompt_template["template"]), (
|
||||
"`prompt_template['template']` must contain a placeholder `{}` for the input text, "
|
||||
f"got {self.prompt_template['template']}")
|
||||
f"got {self.prompt_template['template']}"
|
||||
)
|
||||
|
||||
self.use_video_template = self.prompt_template_video is not None
|
||||
if self.use_video_template:
|
||||
if self.prompt_template_video is not None:
|
||||
assert (
|
||||
isinstance(self.prompt_template_video, dict) and "template" in self.prompt_template_video
|
||||
isinstance(self.prompt_template_video, dict)
|
||||
and "template" in self.prompt_template_video
|
||||
), f"`prompt_template_video` must be a dictionary with a key 'template', got {self.prompt_template_video}"
|
||||
assert "{}" in str(self.prompt_template_video["template"]), (
|
||||
"`prompt_template_video['template']` must contain a placeholder `{}` for the input text, "
|
||||
f"got {self.prompt_template_video['template']}")
|
||||
f"got {self.prompt_template_video['template']}"
|
||||
)
|
||||
|
||||
if "t5" in text_encoder_type:
|
||||
self.output_key = output_key or "last_hidden_state"
|
||||
@@ -183,7 +205,7 @@ class TextEncoder(nn.Module):
|
||||
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
|
||||
prevent_empty_text (bool): If Ture, we will prevent the user text from being empty
|
||||
by adding a space. Defaults to True.
|
||||
"""
|
||||
if isinstance(template, str):
|
||||
@@ -208,7 +230,10 @@ class TextEncoder(nn.Module):
|
||||
else:
|
||||
raise ValueError(f"Unsupported data type: {data_type}")
|
||||
if isinstance(text, (list, tuple)):
|
||||
text = [self.apply_text_to_template(one_text, prompt_template) for one_text in text]
|
||||
text = [
|
||||
self.apply_text_to_template(one_text, prompt_template)
|
||||
for one_text in text
|
||||
]
|
||||
if isinstance(text[0], list):
|
||||
tokenize_input_type = "list"
|
||||
elif isinstance(text, str):
|
||||
@@ -270,13 +295,18 @@ class TextEncoder(nn.Module):
|
||||
"""
|
||||
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)
|
||||
hidden_state_skip_layer = use_default(
|
||||
hidden_state_skip_layer, self.hidden_state_skip_layer
|
||||
)
|
||||
do_sample = use_default(do_sample, not self.reproduce)
|
||||
attention_mask = (batch_encoding["attention_mask"].to(device) if use_attention_mask else None)
|
||||
attention_mask = (
|
||||
batch_encoding["attention_mask"].to(device) if use_attention_mask else None
|
||||
)
|
||||
outputs = self.model(
|
||||
input_ids=batch_encoding["input_ids"].to(device),
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=output_hidden_states or hidden_state_skip_layer is not None,
|
||||
output_hidden_states=output_hidden_states
|
||||
or 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)]
|
||||
@@ -297,10 +327,14 @@ class TextEncoder(nn.Module):
|
||||
raise ValueError(f"Unsupported data type: {data_type}")
|
||||
if crop_start > 0:
|
||||
last_hidden_state = last_hidden_state[:, crop_start:]
|
||||
attention_mask = (attention_mask[:, crop_start:] if use_attention_mask else None)
|
||||
attention_mask = (
|
||||
attention_mask[:, crop_start:] if use_attention_mask else None
|
||||
)
|
||||
|
||||
if output_hidden_states:
|
||||
return TextEncoderModelOutput(last_hidden_state, attention_mask, outputs.hidden_states)
|
||||
return TextEncoderModelOutput(
|
||||
last_hidden_state, attention_mask, outputs.hidden_states
|
||||
)
|
||||
return TextEncoderModelOutput(last_hidden_state, attention_mask)
|
||||
|
||||
def forward(
|
||||
|
||||
@@ -1,8 +1,9 @@
|
||||
import numpy as np
|
||||
import math
|
||||
|
||||
|
||||
def align_to(value, alignment):
|
||||
"""align height, width according to alignment
|
||||
"""align hight, width according to alignment
|
||||
|
||||
Args:
|
||||
value (int): height or width
|
||||
|
||||
@@ -1,11 +1,11 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
from einops import rearrange
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
import numpy as np
|
||||
import imageio
|
||||
|
||||
CODE_SUFFIXES = {
|
||||
".py", # Python codes
|
||||
|
||||
@@ -1,9 +1,9 @@
|
||||
import collections.abc
|
||||
|
||||
from itertools import repeat
|
||||
|
||||
|
||||
def _ntuple(n):
|
||||
|
||||
def parse(x):
|
||||
if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):
|
||||
x = tuple(x)
|
||||
@@ -25,7 +25,7 @@ def as_tuple(x):
|
||||
if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):
|
||||
return tuple(x)
|
||||
if x is None or isinstance(x, (int, float, str)):
|
||||
return (x, )
|
||||
return (x,)
|
||||
else:
|
||||
raise ValueError(f"Unknown type {type(x)}")
|
||||
|
||||
|
||||
@@ -1,16 +1,16 @@
|
||||
import argparse
|
||||
|
||||
import torch
|
||||
from transformers import AutoProcessor, LlavaForConditionalGeneration
|
||||
from transformers import (
|
||||
AutoProcessor,
|
||||
LlavaForConditionalGeneration,
|
||||
)
|
||||
|
||||
|
||||
def preprocess_text_encoder_tokenizer(args):
|
||||
|
||||
processor = AutoProcessor.from_pretrained(args.input_dir)
|
||||
model = LlavaForConditionalGeneration.from_pretrained(
|
||||
args.input_dir,
|
||||
torch_dtype=torch.float16,
|
||||
low_cpu_mem_usage=True,
|
||||
args.input_dir, torch_dtype=torch.float16, low_cpu_mem_usage=True,
|
||||
).to(0)
|
||||
|
||||
model.language_model.save_pretrained(f"{args.output_dir}")
|
||||
|
||||
@@ -2,8 +2,8 @@ from pathlib import Path
|
||||
|
||||
import torch
|
||||
|
||||
from ..constants import PRECISION_TO_TYPE, VAE_PATH
|
||||
from .autoencoder_kl_causal_3d import AutoencoderKLCausal3D
|
||||
from ..constants import VAE_PATH, PRECISION_TO_TYPE
|
||||
|
||||
|
||||
def load_vae(
|
||||
@@ -14,7 +14,7 @@ def load_vae(
|
||||
logger=None,
|
||||
device=None,
|
||||
):
|
||||
"""the function to load the 3D VAE model
|
||||
"""the fucntion to load the 3D VAE model
|
||||
|
||||
Args:
|
||||
vae_type (str): the type of the 3D VAE model. Defaults to "884-16c-hy".
|
||||
@@ -42,7 +42,9 @@ def load_vae(
|
||||
if "state_dict" in ckpt:
|
||||
ckpt = ckpt["state_dict"]
|
||||
if any(k.startswith("vae.") for k in ckpt.keys()):
|
||||
ckpt = {k.replace("vae.", ""): v for k, v in ckpt.items() if k.startswith("vae.")}
|
||||
ckpt = {
|
||||
k.replace("vae.", ""): v for k, v in ckpt.items() if k.startswith("vae.")
|
||||
}
|
||||
vae.load_state_dict(ckpt)
|
||||
|
||||
spatial_compression_ratio = vae.config.spatial_compression_ratio
|
||||
|
||||
@@ -16,16 +16,13 @@
|
||||
# Modified from diffusers==0.29.2
|
||||
#
|
||||
# ==============================================================================
|
||||
from dataclasses import dataclass
|
||||
from math import prod
|
||||
from typing import Dict, Optional, Tuple, Union
|
||||
from dataclasses import dataclass
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
|
||||
from fastvideo.utils.parallel_states import nccl_info
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
|
||||
try:
|
||||
# This diffusers is modified and packed in the mirror.
|
||||
@@ -33,15 +30,26 @@ try:
|
||||
except ImportError:
|
||||
# Use this to be compatible with the original diffusers.
|
||||
from diffusers.loaders.single_file_model import (
|
||||
FromOriginalModelMixin as FromOriginalVAEMixin, )
|
||||
|
||||
from diffusers.models.attention_processor import (ADDED_KV_ATTENTION_PROCESSORS, CROSS_ATTENTION_PROCESSORS, Attention,
|
||||
AttentionProcessor, AttnAddedKVProcessor, AttnProcessor)
|
||||
FromOriginalModelMixin as FromOriginalVAEMixin,
|
||||
)
|
||||
from diffusers.utils.accelerate_utils import apply_forward_hook
|
||||
from diffusers.models.attention_processor import (
|
||||
ADDED_KV_ATTENTION_PROCESSORS,
|
||||
CROSS_ATTENTION_PROCESSORS,
|
||||
Attention,
|
||||
AttentionProcessor,
|
||||
AttnAddedKVProcessor,
|
||||
AttnProcessor,
|
||||
)
|
||||
from diffusers.models.modeling_outputs import AutoencoderKLOutput
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.utils.accelerate_utils import apply_forward_hook
|
||||
|
||||
from .vae import BaseOutput, DecoderCausal3D, DecoderOutput, DiagonalGaussianDistribution, EncoderCausal3D
|
||||
from .vae import (
|
||||
DecoderCausal3D,
|
||||
BaseOutput,
|
||||
DecoderOutput,
|
||||
DiagonalGaussianDistribution,
|
||||
EncoderCausal3D,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -65,9 +73,9 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
self,
|
||||
in_channels: int = 3,
|
||||
out_channels: int = 3,
|
||||
down_block_types: Tuple[str] = ("DownEncoderBlockCausal3D", ),
|
||||
up_block_types: Tuple[str] = ("UpDecoderBlockCausal3D", ),
|
||||
block_out_channels: Tuple[int] = (64, ),
|
||||
down_block_types: Tuple[str] = ("DownEncoderBlockCausal3D",),
|
||||
up_block_types: Tuple[str] = ("UpDecoderBlockCausal3D",),
|
||||
block_out_channels: Tuple[int] = (64,),
|
||||
layers_per_block: int = 1,
|
||||
act_fn: str = "silu",
|
||||
latent_channels: int = 4,
|
||||
@@ -111,22 +119,30 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
mid_block_add_attention=mid_block_add_attention,
|
||||
)
|
||||
|
||||
self.quant_conv = nn.Conv3d(2 * latent_channels, 2 * latent_channels, kernel_size=1)
|
||||
self.post_quant_conv = nn.Conv3d(latent_channels, latent_channels, kernel_size=1)
|
||||
self.quant_conv = nn.Conv3d(
|
||||
2 * latent_channels, 2 * latent_channels, kernel_size=1
|
||||
)
|
||||
self.post_quant_conv = nn.Conv3d(
|
||||
latent_channels, latent_channels, kernel_size=1
|
||||
)
|
||||
|
||||
self.use_slicing = False
|
||||
self.use_spatial_tiling = False
|
||||
self.use_temporal_tiling = False
|
||||
self.use_parallel = False
|
||||
|
||||
# only relevant if vae tiling is enabled
|
||||
self.tile_sample_min_tsize = sample_tsize
|
||||
self.tile_latent_min_tsize = sample_tsize // time_compression_ratio
|
||||
|
||||
self.tile_sample_min_size = self.config.sample_size
|
||||
sample_size = (self.config.sample_size[0] if isinstance(self.config.sample_size,
|
||||
(list, tuple)) else self.config.sample_size)
|
||||
self.tile_latent_min_size = int(sample_size / (2**(len(self.config.block_out_channels) - 1)))
|
||||
sample_size = (
|
||||
self.config.sample_size[0]
|
||||
if isinstance(self.config.sample_size, (list, tuple))
|
||||
else self.config.sample_size
|
||||
)
|
||||
self.tile_latent_min_size = int(
|
||||
sample_size / (2 ** (len(self.config.block_out_channels) - 1))
|
||||
)
|
||||
self.tile_overlap_factor = 0.25
|
||||
|
||||
def _set_gradient_checkpointing(self, module, value=False):
|
||||
@@ -162,12 +178,6 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
self.disable_spatial_tiling()
|
||||
self.disable_temporal_tiling()
|
||||
|
||||
def enable_parallel(self):
|
||||
r"""
|
||||
Enable sequence parallelism for the model. This will allow the vae to decode (with tiling) in parallel.
|
||||
"""
|
||||
self.use_parallel = True
|
||||
|
||||
def enable_slicing(self):
|
||||
r"""
|
||||
Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to
|
||||
@@ -199,7 +209,9 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
processors: Dict[str, AttentionProcessor],
|
||||
):
|
||||
if hasattr(module, "get_processor"):
|
||||
processors[f"{name}.processor"] = module.get_processor(return_deprecated_lora=True)
|
||||
processors[f"{name}.processor"] = module.get_processor(
|
||||
return_deprecated_lora=True
|
||||
)
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
|
||||
@@ -234,14 +246,17 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
if isinstance(processor, dict) and len(processor) != count:
|
||||
raise ValueError(
|
||||
f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
|
||||
f" number of attention layers: {count}. Please make sure to pass {count} processor classes.")
|
||||
f" number of attention layers: {count}. Please make sure to pass {count} processor classes."
|
||||
)
|
||||
|
||||
def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
|
||||
if hasattr(module, "set_processor"):
|
||||
if not isinstance(processor, dict):
|
||||
module.set_processor(processor, _remove_lora=_remove_lora)
|
||||
else:
|
||||
module.set_processor(processor.pop(f"{name}.processor"), _remove_lora=_remove_lora)
|
||||
module.set_processor(
|
||||
processor.pop(f"{name}.processor"), _remove_lora=_remove_lora
|
||||
)
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
|
||||
@@ -254,9 +269,15 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
"""
|
||||
Disables custom attention processors and sets the default attention implementation.
|
||||
"""
|
||||
if all(proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS for proc in self.attn_processors.values()):
|
||||
if all(
|
||||
proc.__class__ in ADDED_KV_ATTENTION_PROCESSORS
|
||||
for proc in self.attn_processors.values()
|
||||
):
|
||||
processor = AttnAddedKVProcessor()
|
||||
elif all(proc.__class__ in CROSS_ATTENTION_PROCESSORS for proc in self.attn_processors.values()):
|
||||
elif all(
|
||||
proc.__class__ in CROSS_ATTENTION_PROCESSORS
|
||||
for proc in self.attn_processors.values()
|
||||
):
|
||||
processor = AttnProcessor()
|
||||
else:
|
||||
raise ValueError(
|
||||
@@ -266,9 +287,9 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
self.set_attn_processor(processor, _remove_lora=True)
|
||||
|
||||
@apply_forward_hook
|
||||
def encode(self,
|
||||
x: torch.FloatTensor,
|
||||
return_dict: bool = True) -> Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]:
|
||||
def encode(
|
||||
self, x: torch.FloatTensor, return_dict: bool = True
|
||||
) -> Union[AutoencoderKLOutput, Tuple[DiagonalGaussianDistribution]]:
|
||||
"""
|
||||
Encode a batch of images/videos into latents.
|
||||
|
||||
@@ -286,8 +307,10 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
if self.use_temporal_tiling and x.shape[2] > self.tile_sample_min_tsize:
|
||||
return self.temporal_tiled_encode(x, return_dict=return_dict)
|
||||
|
||||
if self.use_spatial_tiling and (x.shape[-1] > self.tile_sample_min_size
|
||||
or x.shape[-2] > self.tile_sample_min_size):
|
||||
if self.use_spatial_tiling and (
|
||||
x.shape[-1] > self.tile_sample_min_size
|
||||
or x.shape[-2] > self.tile_sample_min_size
|
||||
):
|
||||
return self.spatial_tiled_encode(x, return_dict=return_dict)
|
||||
|
||||
if self.use_slicing and x.shape[0] > 1:
|
||||
@@ -300,36 +323,36 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
|
||||
if not return_dict:
|
||||
return (posterior, )
|
||||
return (posterior,)
|
||||
|
||||
return AutoencoderKLOutput(latent_dist=posterior)
|
||||
|
||||
def _decode(self, z: torch.FloatTensor, return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
def _decode(
|
||||
self, z: torch.FloatTensor, return_dict: bool = True
|
||||
) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
assert len(z.shape) == 5, "The input tensor should have 5 dimensions."
|
||||
|
||||
if self.use_parallel:
|
||||
return self.parallel_tiled_decode(z, return_dict=return_dict)
|
||||
|
||||
if self.use_temporal_tiling and z.shape[2] > self.tile_latent_min_tsize:
|
||||
return self.temporal_tiled_decode(z, return_dict=return_dict)
|
||||
|
||||
if self.use_spatial_tiling and (z.shape[-1] > self.tile_latent_min_size
|
||||
or z.shape[-2] > self.tile_latent_min_size):
|
||||
if self.use_spatial_tiling and (
|
||||
z.shape[-1] > self.tile_latent_min_size
|
||||
or z.shape[-2] > self.tile_latent_min_size
|
||||
):
|
||||
return self.spatial_tiled_decode(z, return_dict=return_dict)
|
||||
|
||||
z = self.post_quant_conv(z)
|
||||
dec = self.decoder(z)
|
||||
|
||||
if not return_dict:
|
||||
return (dec, )
|
||||
return (dec,)
|
||||
|
||||
return DecoderOutput(sample=dec)
|
||||
|
||||
@apply_forward_hook
|
||||
def decode(self,
|
||||
z: torch.FloatTensor,
|
||||
return_dict: bool = True,
|
||||
generator=None) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
def decode(
|
||||
self, z: torch.FloatTensor, return_dict: bool = True, generator=None
|
||||
) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
"""
|
||||
Decode a batch of images/videos.
|
||||
|
||||
@@ -351,30 +374,38 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
decoded = self._decode(z).sample
|
||||
|
||||
if not return_dict:
|
||||
return (decoded, )
|
||||
return (decoded,)
|
||||
|
||||
return DecoderOutput(sample=decoded)
|
||||
|
||||
def blend_v(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
|
||||
def blend_v(
|
||||
self, a: torch.Tensor, b: torch.Tensor, blend_extent: int
|
||||
) -> torch.Tensor:
|
||||
blend_extent = min(a.shape[-2], b.shape[-2], blend_extent)
|
||||
for y in range(blend_extent):
|
||||
b[:, :, :,
|
||||
y, :] = a[:, :, :, -blend_extent + y, :] * (1 - y / blend_extent) + b[:, :, :, y, :] * (y / blend_extent)
|
||||
b[:, :, :, y, :] = a[:, :, :, -blend_extent + y, :] * (
|
||||
1 - y / blend_extent
|
||||
) + b[:, :, :, y, :] * (y / blend_extent)
|
||||
return b
|
||||
|
||||
def blend_h(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
|
||||
def blend_h(
|
||||
self, a: torch.Tensor, b: torch.Tensor, blend_extent: int
|
||||
) -> torch.Tensor:
|
||||
blend_extent = min(a.shape[-1], b.shape[-1], blend_extent)
|
||||
for x in range(blend_extent):
|
||||
b[:, :, :, :,
|
||||
x] = a[:, :, :, :, -blend_extent + x] * (1 - x / blend_extent) + b[:, :, :, :, x] * (x / blend_extent)
|
||||
b[:, :, :, :, x] = a[:, :, :, :, -blend_extent + x] * (
|
||||
1 - x / blend_extent
|
||||
) + b[:, :, :, :, x] * (x / blend_extent)
|
||||
return b
|
||||
|
||||
def blend_t(self, a: torch.Tensor, b: torch.Tensor, blend_extent: int) -> torch.Tensor:
|
||||
def blend_t(
|
||||
self, a: torch.Tensor, b: torch.Tensor, blend_extent: int
|
||||
) -> torch.Tensor:
|
||||
blend_extent = min(a.shape[-3], b.shape[-3], blend_extent)
|
||||
for x in range(blend_extent):
|
||||
b[:, :,
|
||||
x, :, :] = a[:, :, -blend_extent + x, :, :] * (1 - x / blend_extent) + b[:, :,
|
||||
x, :, :] * (x / blend_extent)
|
||||
b[:, :, x, :, :] = a[:, :, -blend_extent + x, :, :] * (
|
||||
1 - x / blend_extent
|
||||
) + b[:, :, x, :, :] * (x / blend_extent)
|
||||
return b
|
||||
|
||||
def spatial_tiled_encode(
|
||||
@@ -410,7 +441,13 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
for i in range(0, x.shape[-2], overlap_size):
|
||||
row = []
|
||||
for j in range(0, x.shape[-1], overlap_size):
|
||||
tile = x[:, :, :, i:i + self.tile_sample_min_size, j:j + self.tile_sample_min_size, ]
|
||||
tile = x[
|
||||
:,
|
||||
:,
|
||||
:,
|
||||
i : i + self.tile_sample_min_size,
|
||||
j : j + self.tile_sample_min_size,
|
||||
]
|
||||
tile = self.encoder(tile)
|
||||
tile = self.quant_conv(tile)
|
||||
row.append(tile)
|
||||
@@ -434,13 +471,13 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
if not return_dict:
|
||||
return (posterior, )
|
||||
return (posterior,)
|
||||
|
||||
return AutoencoderKLOutput(latent_dist=posterior)
|
||||
|
||||
def spatial_tiled_decode(self,
|
||||
z: torch.FloatTensor,
|
||||
return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
def spatial_tiled_decode(
|
||||
self, z: torch.FloatTensor, return_dict: bool = True
|
||||
) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
r"""
|
||||
Decode a batch of images/videos using a tiled decoder.
|
||||
|
||||
@@ -464,7 +501,13 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
for i in range(0, z.shape[-2], overlap_size):
|
||||
row = []
|
||||
for j in range(0, z.shape[-1], overlap_size):
|
||||
tile = z[:, :, :, i:i + self.tile_latent_min_size, j:j + self.tile_latent_min_size, ]
|
||||
tile = z[
|
||||
:,
|
||||
:,
|
||||
:,
|
||||
i : i + self.tile_latent_min_size,
|
||||
j : j + self.tile_latent_min_size,
|
||||
]
|
||||
tile = self.post_quant_conv(tile)
|
||||
decoded = self.decoder(tile)
|
||||
row.append(decoded)
|
||||
@@ -484,11 +527,13 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
dec = torch.cat(result_rows, dim=-2)
|
||||
if not return_dict:
|
||||
return (dec, )
|
||||
return (dec,)
|
||||
|
||||
return DecoderOutput(sample=dec)
|
||||
|
||||
def temporal_tiled_encode(self, x: torch.FloatTensor, return_dict: bool = True) -> AutoencoderKLOutput:
|
||||
def temporal_tiled_encode(
|
||||
self, x: torch.FloatTensor, return_dict: bool = True
|
||||
) -> AutoencoderKLOutput:
|
||||
|
||||
B, C, T, H, W = x.shape
|
||||
overlap_size = int(self.tile_sample_min_tsize * (1 - self.tile_overlap_factor))
|
||||
@@ -498,9 +543,11 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
# Split the video into tiles and encode them separately.
|
||||
row = []
|
||||
for i in range(0, T, overlap_size):
|
||||
tile = x[:, :, i:i + self.tile_sample_min_tsize + 1, :, :]
|
||||
if self.use_spatial_tiling and (tile.shape[-1] > self.tile_sample_min_size
|
||||
or tile.shape[-2] > self.tile_sample_min_size):
|
||||
tile = x[:, :, i : i + self.tile_sample_min_tsize + 1, :, :]
|
||||
if self.use_spatial_tiling and (
|
||||
tile.shape[-1] > self.tile_sample_min_size
|
||||
or tile.shape[-2] > self.tile_sample_min_size
|
||||
):
|
||||
tile = self.spatial_tiled_encode(tile, return_moments=True)
|
||||
else:
|
||||
tile = self.encoder(tile)
|
||||
@@ -514,19 +561,19 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
tile = self.blend_t(row[i - 1], tile, blend_extent)
|
||||
result_row.append(tile[:, :, :t_limit, :, :])
|
||||
else:
|
||||
result_row.append(tile[:, :, :t_limit + 1, :, :])
|
||||
result_row.append(tile[:, :, : t_limit + 1, :, :])
|
||||
|
||||
moments = torch.cat(result_row, dim=2)
|
||||
posterior = DiagonalGaussianDistribution(moments)
|
||||
|
||||
if not return_dict:
|
||||
return (posterior, )
|
||||
return (posterior,)
|
||||
|
||||
return AutoencoderKLOutput(latent_dist=posterior)
|
||||
|
||||
def temporal_tiled_decode(self,
|
||||
z: torch.FloatTensor,
|
||||
return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
def temporal_tiled_decode(
|
||||
self, z: torch.FloatTensor, return_dict: bool = True
|
||||
) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
# Split z into overlapping tiles and decode them separately.
|
||||
|
||||
B, C, T, H, W = z.shape
|
||||
@@ -536,9 +583,11 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
row = []
|
||||
for i in range(0, T, overlap_size):
|
||||
tile = z[:, :, i:i + self.tile_latent_min_tsize + 1, :, :]
|
||||
if self.use_spatial_tiling and (tile.shape[-1] > self.tile_latent_min_size
|
||||
or tile.shape[-2] > self.tile_latent_min_size):
|
||||
tile = z[:, :, i : i + self.tile_latent_min_tsize + 1, :, :]
|
||||
if self.use_spatial_tiling and (
|
||||
tile.shape[-1] > self.tile_latent_min_size
|
||||
or tile.shape[-2] > self.tile_latent_min_size
|
||||
):
|
||||
decoded = self.spatial_tiled_decode(tile, return_dict=True).sample
|
||||
else:
|
||||
tile = self.post_quant_conv(tile)
|
||||
@@ -552,145 +601,14 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
tile = self.blend_t(row[i - 1], tile, blend_extent)
|
||||
result_row.append(tile[:, :, :t_limit, :, :])
|
||||
else:
|
||||
result_row.append(tile[:, :, :t_limit + 1, :, :])
|
||||
result_row.append(tile[:, :, : t_limit + 1, :, :])
|
||||
|
||||
dec = torch.cat(result_row, dim=2)
|
||||
if not return_dict:
|
||||
return (dec, )
|
||||
return (dec,)
|
||||
|
||||
return DecoderOutput(sample=dec)
|
||||
|
||||
def _parallel_data_generator(self, gathered_results, gathered_dim_metadata):
|
||||
global_idx = 0
|
||||
for i, per_rank_metadata in enumerate(gathered_dim_metadata):
|
||||
_start_shape = 0
|
||||
for shape in per_rank_metadata:
|
||||
mul_shape = prod(shape)
|
||||
yield (gathered_results[i, _start_shape:_start_shape + mul_shape].reshape(shape), global_idx)
|
||||
_start_shape += mul_shape
|
||||
global_idx += 1
|
||||
|
||||
def parallel_tiled_decode(self,
|
||||
z: torch.FloatTensor,
|
||||
return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
"""
|
||||
Parallel version of tiled_decode that distributes both temporal and spatial computation across GPUs
|
||||
"""
|
||||
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
|
||||
B, C, T, H, W = z.shape
|
||||
|
||||
# Calculate parameters
|
||||
t_overlap_size = int(self.tile_latent_min_tsize * (1 - self.tile_overlap_factor))
|
||||
t_blend_extent = int(self.tile_sample_min_tsize * self.tile_overlap_factor)
|
||||
t_limit = self.tile_sample_min_tsize - t_blend_extent
|
||||
|
||||
s_overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor))
|
||||
s_blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor)
|
||||
s_row_limit = self.tile_sample_min_size - s_blend_extent
|
||||
|
||||
# Calculate tile dimensions
|
||||
num_t_tiles = (T + t_overlap_size - 1) // t_overlap_size
|
||||
num_h_tiles = (H + s_overlap_size - 1) // s_overlap_size
|
||||
num_w_tiles = (W + s_overlap_size - 1) // s_overlap_size
|
||||
total_spatial_tiles = num_h_tiles * num_w_tiles
|
||||
total_tiles = num_t_tiles * total_spatial_tiles
|
||||
|
||||
# Calculate tiles per rank and padding
|
||||
tiles_per_rank = (total_tiles + world_size - 1) // world_size
|
||||
start_tile_idx = rank * tiles_per_rank
|
||||
end_tile_idx = min((rank + 1) * tiles_per_rank, total_tiles)
|
||||
|
||||
local_results = []
|
||||
local_dim_metadata = []
|
||||
# Process assigned tiles
|
||||
for local_idx, global_idx in enumerate(range(start_tile_idx, end_tile_idx)):
|
||||
# Convert flat index to 3D indices
|
||||
t_idx = global_idx // total_spatial_tiles
|
||||
spatial_idx = global_idx % total_spatial_tiles
|
||||
h_idx = spatial_idx // num_w_tiles
|
||||
w_idx = spatial_idx % num_w_tiles
|
||||
|
||||
# Calculate positions
|
||||
t_start = t_idx * t_overlap_size
|
||||
h_start = h_idx * s_overlap_size
|
||||
w_start = w_idx * s_overlap_size
|
||||
|
||||
# Extract and process tile
|
||||
tile = z[:, :, t_start:t_start + self.tile_latent_min_tsize + 1,
|
||||
h_start:h_start + self.tile_latent_min_size, w_start:w_start + self.tile_latent_min_size]
|
||||
|
||||
# Process tile
|
||||
tile = self.post_quant_conv(tile)
|
||||
decoded = self.decoder(tile)
|
||||
|
||||
if t_start > 0:
|
||||
decoded = decoded[:, :, 1:, :, :]
|
||||
|
||||
# Store metadata
|
||||
shape = decoded.shape
|
||||
# Store decoded data (flattened)
|
||||
decoded_flat = decoded.reshape(-1)
|
||||
local_results.append(decoded_flat)
|
||||
local_dim_metadata.append(shape)
|
||||
|
||||
results = torch.cat(local_results, dim=0).contiguous()
|
||||
del local_results
|
||||
torch.cuda.empty_cache()
|
||||
# first gather size to pad the results
|
||||
local_size = torch.tensor([results.size(0)], device=results.device, dtype=torch.int64)
|
||||
all_sizes = [torch.zeros(1, device=results.device, dtype=torch.int64) for _ in range(world_size)]
|
||||
dist.all_gather(all_sizes, local_size)
|
||||
max_size = max(size.item() for size in all_sizes)
|
||||
padded_results = torch.zeros(max_size, device=results.device)
|
||||
padded_results[:results.size(0)] = results
|
||||
del results
|
||||
torch.cuda.empty_cache()
|
||||
# Gather all results
|
||||
gathered_dim_metadata = [None] * world_size
|
||||
gathered_results = torch.zeros_like(padded_results).repeat(
|
||||
world_size, *[1] * len(padded_results.shape)).contiguous(
|
||||
) # use contiguous to make sure it won't copy data in the following operations
|
||||
dist.all_gather_into_tensor(gathered_results, padded_results)
|
||||
dist.all_gather_object(gathered_dim_metadata, local_dim_metadata)
|
||||
# Process gathered results
|
||||
data = [[[[] for _ in range(num_w_tiles)] for _ in range(num_h_tiles)] for _ in range(num_t_tiles)]
|
||||
for current_data, global_idx in self._parallel_data_generator(gathered_results, gathered_dim_metadata):
|
||||
t_idx = global_idx // total_spatial_tiles
|
||||
spatial_idx = global_idx % total_spatial_tiles
|
||||
h_idx = spatial_idx // num_w_tiles
|
||||
w_idx = spatial_idx % num_w_tiles
|
||||
data[t_idx][h_idx][w_idx] = current_data
|
||||
# Merge results
|
||||
result_slices = []
|
||||
last_slice_data = None
|
||||
for i, tem_data in enumerate(data):
|
||||
slice_data = self._merge_spatial_tiles(tem_data, s_blend_extent, s_row_limit)
|
||||
if i > 0:
|
||||
slice_data = self.blend_t(last_slice_data, slice_data, t_blend_extent)
|
||||
result_slices.append(slice_data[:, :, :t_limit, :, :])
|
||||
else:
|
||||
result_slices.append(slice_data[:, :, :t_limit + 1, :, :])
|
||||
last_slice_data = slice_data
|
||||
dec = torch.cat(result_slices, dim=2)
|
||||
|
||||
if not return_dict:
|
||||
return (dec, )
|
||||
return DecoderOutput(sample=dec)
|
||||
|
||||
def _merge_spatial_tiles(self, spatial_rows, blend_extent, row_limit):
|
||||
"""Helper function to merge spatial tiles with blending"""
|
||||
result_rows = []
|
||||
for i, row in enumerate(spatial_rows):
|
||||
result_row = []
|
||||
for j, tile in enumerate(row):
|
||||
if i > 0:
|
||||
tile = self.blend_v(spatial_rows[i - 1][j], tile, blend_extent)
|
||||
if j > 0:
|
||||
tile = self.blend_h(row[j - 1], tile, blend_extent)
|
||||
result_row.append(tile[:, :, :, :row_limit, :row_limit])
|
||||
result_rows.append(torch.cat(result_row, dim=-1))
|
||||
return torch.cat(result_rows, dim=-2)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
sample: torch.FloatTensor,
|
||||
@@ -719,7 +637,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
if return_posterior:
|
||||
return (dec, posterior)
|
||||
else:
|
||||
return (dec, )
|
||||
return (dec,)
|
||||
if return_posterior:
|
||||
return DecoderOutput2(sample=dec, posterior=posterior)
|
||||
else:
|
||||
@@ -741,7 +659,9 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
for _, attn_processor in self.attn_processors.items():
|
||||
if "Added" in str(attn_processor.__class__.__name__):
|
||||
raise ValueError("`fuse_qkv_projections()` is not supported for models having added KV projections.")
|
||||
raise ValueError(
|
||||
"`fuse_qkv_projections()` is not supported for models having added KV projections."
|
||||
)
|
||||
|
||||
self.original_attn_processors = self.attn_processors
|
||||
|
||||
|
||||
@@ -21,22 +21,27 @@ from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from diffusers.models.activations import get_activation
|
||||
from diffusers.models.attention_processor import Attention, SpatialNorm
|
||||
from diffusers.models.normalization import AdaGroupNorm, RMSNorm
|
||||
from diffusers.utils import logging
|
||||
from einops import rearrange
|
||||
from torch import nn
|
||||
from einops import rearrange
|
||||
|
||||
from diffusers.utils import logging
|
||||
from diffusers.models.activations import get_activation
|
||||
from diffusers.models.attention_processor import SpatialNorm
|
||||
from diffusers.models.attention_processor import Attention
|
||||
from diffusers.models.normalization import AdaGroupNorm
|
||||
from diffusers.models.normalization import RMSNorm
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
def prepare_causal_attention_mask(n_frame: int, n_hw: int, dtype, device, batch_size: int = None):
|
||||
def prepare_causal_attention_mask(
|
||||
n_frame: int, n_hw: int, dtype, device, batch_size: int = None
|
||||
):
|
||||
seq_len = n_frame * n_hw
|
||||
mask = torch.full((seq_len, seq_len), float("-inf"), dtype=dtype, device=device)
|
||||
for i in range(seq_len):
|
||||
i_frame = i // n_hw
|
||||
mask[i, :(i_frame + 1) * n_hw] = 0
|
||||
mask[i, : (i_frame + 1) * n_hw] = 0
|
||||
if batch_size is not None:
|
||||
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
|
||||
return mask
|
||||
@@ -71,7 +76,9 @@ class CausalConv3d(nn.Module):
|
||||
) # W, H, T
|
||||
self.time_causal_padding = padding
|
||||
|
||||
self.conv = nn.Conv3d(chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs)
|
||||
self.conv = nn.Conv3d(
|
||||
chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs
|
||||
)
|
||||
|
||||
def forward(self, x):
|
||||
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
|
||||
@@ -84,20 +91,20 @@ class UpsampleCausal3D(nn.Module):
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
channels: int,
|
||||
use_conv: bool = False,
|
||||
use_conv_transpose: bool = False,
|
||||
out_channels: Optional[int] = None,
|
||||
name: str = "conv",
|
||||
kernel_size: Optional[int] = None,
|
||||
padding=1,
|
||||
norm_type=None,
|
||||
eps=None,
|
||||
elementwise_affine=None,
|
||||
bias=True,
|
||||
interpolate=True,
|
||||
upsample_factor=(2, 2, 2),
|
||||
self,
|
||||
channels: int,
|
||||
use_conv: bool = False,
|
||||
use_conv_transpose: bool = False,
|
||||
out_channels: Optional[int] = None,
|
||||
name: str = "conv",
|
||||
kernel_size: Optional[int] = None,
|
||||
padding=1,
|
||||
norm_type=None,
|
||||
eps=None,
|
||||
elementwise_affine=None,
|
||||
bias=True,
|
||||
interpolate=True,
|
||||
upsample_factor=(2, 2, 2),
|
||||
):
|
||||
super().__init__()
|
||||
self.channels = channels
|
||||
@@ -123,7 +130,9 @@ class UpsampleCausal3D(nn.Module):
|
||||
elif use_conv:
|
||||
if kernel_size is None:
|
||||
kernel_size = 3
|
||||
conv = CausalConv3d(self.channels, self.out_channels, kernel_size=kernel_size, bias=bias)
|
||||
conv = CausalConv3d(
|
||||
self.channels, self.out_channels, kernel_size=kernel_size, bias=bias
|
||||
)
|
||||
|
||||
if name == "conv":
|
||||
self.conv = conv
|
||||
@@ -160,10 +169,14 @@ class UpsampleCausal3D(nn.Module):
|
||||
first_h, other_h = hidden_states.split((1, T - 1), dim=2)
|
||||
if output_size is None:
|
||||
if T > 1:
|
||||
other_h = F.interpolate(other_h, scale_factor=self.upsample_factor, mode="nearest")
|
||||
other_h = F.interpolate(
|
||||
other_h, scale_factor=self.upsample_factor, mode="nearest"
|
||||
)
|
||||
|
||||
first_h = first_h.squeeze(2)
|
||||
first_h = F.interpolate(first_h, scale_factor=self.upsample_factor[1:], mode="nearest")
|
||||
first_h = F.interpolate(
|
||||
first_h, scale_factor=self.upsample_factor[1:], mode="nearest"
|
||||
)
|
||||
first_h = first_h.unsqueeze(2)
|
||||
else:
|
||||
raise NotImplementedError
|
||||
@@ -241,11 +254,15 @@ class DownsampleCausal3D(nn.Module):
|
||||
else:
|
||||
self.conv = conv
|
||||
|
||||
def forward(self, hidden_states: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
|
||||
def forward(
|
||||
self, hidden_states: torch.FloatTensor, scale: float = 1.0
|
||||
) -> torch.FloatTensor:
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
|
||||
if self.norm is not None:
|
||||
hidden_states = self.norm(hidden_states.permute(0, 2, 3, 1)).permute(0, 3, 1, 2)
|
||||
hidden_states = self.norm(hidden_states.permute(0, 2, 3, 1)).permute(
|
||||
0, 3, 1, 2
|
||||
)
|
||||
|
||||
assert hidden_states.shape[1] == self.channels
|
||||
|
||||
@@ -306,7 +323,9 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
elif self.time_embedding_norm == "spatial":
|
||||
self.norm1 = SpatialNorm(in_channels, temb_channels)
|
||||
else:
|
||||
self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
|
||||
self.norm1 = torch.nn.GroupNorm(
|
||||
num_groups=groups, num_channels=in_channels, eps=eps, affine=True
|
||||
)
|
||||
|
||||
self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, stride=1)
|
||||
|
||||
@@ -315,10 +334,15 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
self.time_emb_proj = linear_cls(temb_channels, out_channels)
|
||||
elif self.time_embedding_norm == "scale_shift":
|
||||
self.time_emb_proj = linear_cls(temb_channels, 2 * out_channels)
|
||||
elif (self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial"):
|
||||
elif (
|
||||
self.time_embedding_norm == "ada_group"
|
||||
or self.time_embedding_norm == "spatial"
|
||||
):
|
||||
self.time_emb_proj = None
|
||||
else:
|
||||
raise ValueError(f"Unknown time_embedding_norm : {self.time_embedding_norm} ")
|
||||
raise ValueError(
|
||||
f"Unknown time_embedding_norm : {self.time_embedding_norm} "
|
||||
)
|
||||
else:
|
||||
self.time_emb_proj = None
|
||||
|
||||
@@ -327,11 +351,15 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
elif self.time_embedding_norm == "spatial":
|
||||
self.norm2 = SpatialNorm(out_channels, temb_channels)
|
||||
else:
|
||||
self.norm2 = torch.nn.GroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
|
||||
self.norm2 = torch.nn.GroupNorm(
|
||||
num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True
|
||||
)
|
||||
|
||||
self.dropout = torch.nn.Dropout(dropout)
|
||||
conv_3d_out_channels = conv_3d_out_channels or out_channels
|
||||
self.conv2 = CausalConv3d(out_channels, conv_3d_out_channels, kernel_size=3, stride=1)
|
||||
self.conv2 = CausalConv3d(
|
||||
out_channels, conv_3d_out_channels, kernel_size=3, stride=1
|
||||
)
|
||||
|
||||
self.nonlinearity = get_activation(non_linearity)
|
||||
|
||||
@@ -341,8 +369,11 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
elif self.down:
|
||||
self.downsample = DownsampleCausal3D(in_channels, use_conv=False, name="op")
|
||||
|
||||
self.use_in_shortcut = (self.in_channels != conv_3d_out_channels
|
||||
if use_in_shortcut is None else use_in_shortcut)
|
||||
self.use_in_shortcut = (
|
||||
self.in_channels != conv_3d_out_channels
|
||||
if use_in_shortcut is None
|
||||
else use_in_shortcut
|
||||
)
|
||||
|
||||
self.conv_shortcut = None
|
||||
if self.use_in_shortcut:
|
||||
@@ -362,7 +393,10 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
) -> torch.FloatTensor:
|
||||
hidden_states = input_tensor
|
||||
|
||||
if (self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial"):
|
||||
if (
|
||||
self.time_embedding_norm == "ada_group"
|
||||
or self.time_embedding_norm == "spatial"
|
||||
):
|
||||
hidden_states = self.norm1(hidden_states, temb)
|
||||
else:
|
||||
hidden_states = self.norm1(hidden_states)
|
||||
@@ -390,7 +424,10 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
if temb is not None and self.time_embedding_norm == "default":
|
||||
hidden_states = hidden_states + temb
|
||||
|
||||
if (self.time_embedding_norm == "ada_group" or self.time_embedding_norm == "spatial"):
|
||||
if (
|
||||
self.time_embedding_norm == "ada_group"
|
||||
or self.time_embedding_norm == "spatial"
|
||||
):
|
||||
hidden_states = self.norm2(hidden_states, temb)
|
||||
else:
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
@@ -447,7 +484,11 @@ def get_down_block3d(
|
||||
)
|
||||
attention_head_dim = num_attention_heads
|
||||
|
||||
down_block_type = (down_block_type[7:] if down_block_type.startswith("UNetRes") else down_block_type)
|
||||
down_block_type = (
|
||||
down_block_type[7:]
|
||||
if down_block_type.startswith("UNetRes")
|
||||
else down_block_type
|
||||
)
|
||||
if down_block_type == "DownEncoderBlockCausal3D":
|
||||
return DownEncoderBlockCausal3D(
|
||||
num_layers=num_layers,
|
||||
@@ -501,7 +542,9 @@ def get_up_block3d(
|
||||
)
|
||||
attention_head_dim = num_attention_heads
|
||||
|
||||
up_block_type = (up_block_type[7:] if up_block_type.startswith("UNetRes") else up_block_type)
|
||||
up_block_type = (
|
||||
up_block_type[7:] if up_block_type.startswith("UNetRes") else up_block_type
|
||||
)
|
||||
if up_block_type == "UpDecoderBlockCausal3D":
|
||||
return UpDecoderBlockCausal3D(
|
||||
num_layers=num_layers,
|
||||
@@ -542,11 +585,15 @@ class UNetMidBlockCausal3D(nn.Module):
|
||||
output_scale_factor: float = 1.0,
|
||||
):
|
||||
super().__init__()
|
||||
resnet_groups = (resnet_groups if resnet_groups is not None else min(in_channels // 4, 32))
|
||||
resnet_groups = (
|
||||
resnet_groups if resnet_groups is not None else min(in_channels // 4, 32)
|
||||
)
|
||||
self.add_attention = add_attention
|
||||
|
||||
if attn_groups is None:
|
||||
attn_groups = (resnet_groups if resnet_time_scale_shift == "default" else None)
|
||||
attn_groups = (
|
||||
resnet_groups if resnet_time_scale_shift == "default" else None
|
||||
)
|
||||
|
||||
# there is always at least one resnet
|
||||
resnets = [
|
||||
@@ -581,12 +628,17 @@ class UNetMidBlockCausal3D(nn.Module):
|
||||
rescale_output_factor=output_scale_factor,
|
||||
eps=resnet_eps,
|
||||
norm_num_groups=attn_groups,
|
||||
spatial_norm_dim=(temb_channels if resnet_time_scale_shift == "spatial" else None),
|
||||
spatial_norm_dim=(
|
||||
temb_channels
|
||||
if resnet_time_scale_shift == "spatial"
|
||||
else None
|
||||
),
|
||||
residual_connection=True,
|
||||
bias=True,
|
||||
upcast_softmax=True,
|
||||
_from_deprecated_attn_block=True,
|
||||
))
|
||||
)
|
||||
)
|
||||
else:
|
||||
attentions.append(None)
|
||||
|
||||
@@ -602,31 +654,35 @@ class UNetMidBlockCausal3D(nn.Module):
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
))
|
||||
)
|
||||
)
|
||||
|
||||
self.attentions = nn.ModuleList(attentions)
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
|
||||
def forward(self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None) -> torch.FloatTensor:
|
||||
def forward(
|
||||
self, hidden_states: torch.FloatTensor, temb: Optional[torch.FloatTensor] = None
|
||||
) -> torch.FloatTensor:
|
||||
hidden_states = self.resnets[0](hidden_states, temb)
|
||||
for attn, resnet in zip(self.attentions, self.resnets[1:]):
|
||||
if attn is not None:
|
||||
B, C, T, H, W = hidden_states.shape
|
||||
hidden_states = rearrange(hidden_states, "b c f h w -> b (f h w) c")
|
||||
attention_mask = prepare_causal_attention_mask(T,
|
||||
H * W,
|
||||
hidden_states.dtype,
|
||||
hidden_states.device,
|
||||
batch_size=B)
|
||||
hidden_states = attn(hidden_states, temb=temb, attention_mask=attention_mask)
|
||||
hidden_states = rearrange(hidden_states, "b (f h w) c -> b c f h w", f=T, h=H, w=W)
|
||||
attention_mask = prepare_causal_attention_mask(
|
||||
T, H * W, hidden_states.dtype, hidden_states.device, batch_size=B
|
||||
)
|
||||
hidden_states = attn(
|
||||
hidden_states, temb=temb, attention_mask=attention_mask
|
||||
)
|
||||
hidden_states = rearrange(
|
||||
hidden_states, "b (f h w) c -> b c f h w", f=T, h=H, w=W
|
||||
)
|
||||
hidden_states = resnet(hidden_states, temb)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class DownEncoderBlockCausal3D(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
@@ -660,25 +716,30 @@ class DownEncoderBlockCausal3D(nn.Module):
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
))
|
||||
)
|
||||
)
|
||||
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
|
||||
if add_downsample:
|
||||
self.downsamplers = nn.ModuleList([
|
||||
DownsampleCausal3D(
|
||||
out_channels,
|
||||
use_conv=True,
|
||||
out_channels=out_channels,
|
||||
padding=downsample_padding,
|
||||
name="op",
|
||||
stride=downsample_stride,
|
||||
)
|
||||
])
|
||||
self.downsamplers = nn.ModuleList(
|
||||
[
|
||||
DownsampleCausal3D(
|
||||
out_channels,
|
||||
use_conv=True,
|
||||
out_channels=out_channels,
|
||||
padding=downsample_padding,
|
||||
name="op",
|
||||
stride=downsample_stride,
|
||||
)
|
||||
]
|
||||
)
|
||||
else:
|
||||
self.downsamplers = None
|
||||
|
||||
def forward(self, hidden_states: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
|
||||
def forward(
|
||||
self, hidden_states: torch.FloatTensor, scale: float = 1.0
|
||||
) -> torch.FloatTensor:
|
||||
for resnet in self.resnets:
|
||||
hidden_states = resnet(hidden_states, temb=None, scale=scale)
|
||||
|
||||
@@ -690,23 +751,22 @@ class DownEncoderBlockCausal3D(nn.Module):
|
||||
|
||||
|
||||
class UpDecoderBlockCausal3D(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
resolution_idx: Optional[int] = None,
|
||||
dropout: float = 0.0,
|
||||
num_layers: int = 1,
|
||||
resnet_eps: float = 1e-6,
|
||||
resnet_time_scale_shift: str = "default", # default, spatial
|
||||
resnet_act_fn: str = "swish",
|
||||
resnet_groups: int = 32,
|
||||
resnet_pre_norm: bool = True,
|
||||
output_scale_factor: float = 1.0,
|
||||
add_upsample: bool = True,
|
||||
upsample_scale_factor=(2, 2, 2),
|
||||
temb_channels: Optional[int] = None,
|
||||
self,
|
||||
in_channels: int,
|
||||
out_channels: int,
|
||||
resolution_idx: Optional[int] = None,
|
||||
dropout: float = 0.0,
|
||||
num_layers: int = 1,
|
||||
resnet_eps: float = 1e-6,
|
||||
resnet_time_scale_shift: str = "default", # default, spatial
|
||||
resnet_act_fn: str = "swish",
|
||||
resnet_groups: int = 32,
|
||||
resnet_pre_norm: bool = True,
|
||||
output_scale_factor: float = 1.0,
|
||||
add_upsample: bool = True,
|
||||
upsample_scale_factor=(2, 2, 2),
|
||||
temb_channels: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
resnets = []
|
||||
@@ -726,19 +786,22 @@ class UpDecoderBlockCausal3D(nn.Module):
|
||||
non_linearity=resnet_act_fn,
|
||||
output_scale_factor=output_scale_factor,
|
||||
pre_norm=resnet_pre_norm,
|
||||
))
|
||||
)
|
||||
)
|
||||
|
||||
self.resnets = nn.ModuleList(resnets)
|
||||
|
||||
if add_upsample:
|
||||
self.upsamplers = nn.ModuleList([
|
||||
UpsampleCausal3D(
|
||||
out_channels,
|
||||
use_conv=True,
|
||||
out_channels=out_channels,
|
||||
upsample_factor=upsample_scale_factor,
|
||||
)
|
||||
])
|
||||
self.upsamplers = nn.ModuleList(
|
||||
[
|
||||
UpsampleCausal3D(
|
||||
out_channels,
|
||||
use_conv=True,
|
||||
out_channels=out_channels,
|
||||
upsample_factor=upsample_scale_factor,
|
||||
)
|
||||
]
|
||||
)
|
||||
else:
|
||||
self.upsamplers = None
|
||||
|
||||
|
||||
@@ -4,11 +4,16 @@ from typing import Optional, Tuple
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from diffusers.models.attention_processor import SpatialNorm
|
||||
|
||||
from diffusers.utils import BaseOutput, is_torch_version
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from .unet_causal_3d_blocks import CausalConv3d, UNetMidBlockCausal3D, get_down_block3d, get_up_block3d
|
||||
from diffusers.models.attention_processor import SpatialNorm
|
||||
from .unet_causal_3d_blocks import (
|
||||
CausalConv3d,
|
||||
UNetMidBlockCausal3D,
|
||||
get_down_block3d,
|
||||
get_up_block3d,
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -33,8 +38,8 @@ class EncoderCausal3D(nn.Module):
|
||||
self,
|
||||
in_channels: int = 3,
|
||||
out_channels: int = 3,
|
||||
down_block_types: Tuple[str, ...] = ("DownEncoderBlockCausal3D", ),
|
||||
block_out_channels: Tuple[int, ...] = (64, ),
|
||||
down_block_types: Tuple[str, ...] = ("DownEncoderBlockCausal3D",),
|
||||
block_out_channels: Tuple[int, ...] = (64,),
|
||||
layers_per_block: int = 2,
|
||||
norm_num_groups: int = 32,
|
||||
act_fn: str = "silu",
|
||||
@@ -46,7 +51,9 @@ class EncoderCausal3D(nn.Module):
|
||||
super().__init__()
|
||||
self.layers_per_block = layers_per_block
|
||||
|
||||
self.conv_in = CausalConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1)
|
||||
self.conv_in = CausalConv3d(
|
||||
in_channels, block_out_channels[0], kernel_size=3, stride=1
|
||||
)
|
||||
self.mid_block = None
|
||||
self.down_blocks = nn.ModuleList([])
|
||||
|
||||
@@ -61,13 +68,17 @@ class EncoderCausal3D(nn.Module):
|
||||
|
||||
if time_compression_ratio == 4:
|
||||
add_spatial_downsample = bool(i < num_spatial_downsample_layers)
|
||||
add_time_downsample = bool(i >= (len(block_out_channels) - 1 - num_time_downsample_layers)
|
||||
and not is_final_block)
|
||||
add_time_downsample = bool(
|
||||
i >= (len(block_out_channels) - 1 - num_time_downsample_layers)
|
||||
and not is_final_block
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported time_compression_ratio: {time_compression_ratio}.")
|
||||
raise ValueError(
|
||||
f"Unsupported time_compression_ratio: {time_compression_ratio}."
|
||||
)
|
||||
|
||||
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
|
||||
downsample_stride_T = (2, ) if add_time_downsample else (1, )
|
||||
downsample_stride_T = (2,) if add_time_downsample else (1,)
|
||||
downsample_stride = tuple(downsample_stride_T + downsample_stride_HW)
|
||||
down_block = get_down_block3d(
|
||||
down_block_type,
|
||||
@@ -99,11 +110,15 @@ class EncoderCausal3D(nn.Module):
|
||||
)
|
||||
|
||||
# out
|
||||
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6)
|
||||
self.conv_norm_out = nn.GroupNorm(
|
||||
num_channels=block_out_channels[-1], num_groups=norm_num_groups, eps=1e-6
|
||||
)
|
||||
self.conv_act = nn.SiLU()
|
||||
|
||||
conv_out_channels = 2 * out_channels if double_z else out_channels
|
||||
self.conv_out = CausalConv3d(block_out_channels[-1], conv_out_channels, kernel_size=3)
|
||||
self.conv_out = CausalConv3d(
|
||||
block_out_channels[-1], conv_out_channels, kernel_size=3
|
||||
)
|
||||
|
||||
def forward(self, sample: torch.FloatTensor) -> torch.FloatTensor:
|
||||
r"""The forward method of the `EncoderCausal3D` class."""
|
||||
@@ -135,8 +150,8 @@ class DecoderCausal3D(nn.Module):
|
||||
self,
|
||||
in_channels: int = 3,
|
||||
out_channels: int = 3,
|
||||
up_block_types: Tuple[str, ...] = ("UpDecoderBlockCausal3D", ),
|
||||
block_out_channels: Tuple[int, ...] = (64, ),
|
||||
up_block_types: Tuple[str, ...] = ("UpDecoderBlockCausal3D",),
|
||||
block_out_channels: Tuple[int, ...] = (64,),
|
||||
layers_per_block: int = 2,
|
||||
norm_num_groups: int = 32,
|
||||
act_fn: str = "silu",
|
||||
@@ -148,7 +163,9 @@ class DecoderCausal3D(nn.Module):
|
||||
super().__init__()
|
||||
self.layers_per_block = layers_per_block
|
||||
|
||||
self.conv_in = CausalConv3d(in_channels, block_out_channels[-1], kernel_size=3, stride=1)
|
||||
self.conv_in = CausalConv3d(
|
||||
in_channels, block_out_channels[-1], kernel_size=3, stride=1
|
||||
)
|
||||
self.mid_block = None
|
||||
self.up_blocks = nn.ModuleList([])
|
||||
|
||||
@@ -179,14 +196,20 @@ class DecoderCausal3D(nn.Module):
|
||||
|
||||
if time_compression_ratio == 4:
|
||||
add_spatial_upsample = bool(i < num_spatial_upsample_layers)
|
||||
add_time_upsample = bool(i >= len(block_out_channels) - 1 - num_time_upsample_layers
|
||||
and not is_final_block)
|
||||
add_time_upsample = bool(
|
||||
i >= len(block_out_channels) - 1 - num_time_upsample_layers
|
||||
and not is_final_block
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unsupported time_compression_ratio: {time_compression_ratio}.")
|
||||
raise ValueError(
|
||||
f"Unsupported time_compression_ratio: {time_compression_ratio}."
|
||||
)
|
||||
|
||||
upsample_scale_factor_HW = (2, 2) if add_spatial_upsample else (1, 1)
|
||||
upsample_scale_factor_T = (2, ) if add_time_upsample else (1, )
|
||||
upsample_scale_factor = tuple(upsample_scale_factor_T + upsample_scale_factor_HW)
|
||||
upsample_scale_factor_T = (2,) if add_time_upsample else (1,)
|
||||
upsample_scale_factor = tuple(
|
||||
upsample_scale_factor_T + upsample_scale_factor_HW
|
||||
)
|
||||
up_block = get_up_block3d(
|
||||
up_block_type,
|
||||
num_layers=self.layers_per_block + 1,
|
||||
@@ -209,7 +232,9 @@ class DecoderCausal3D(nn.Module):
|
||||
if norm_type == "spatial":
|
||||
self.conv_norm_out = SpatialNorm(block_out_channels[0], temb_channels)
|
||||
else:
|
||||
self.conv_norm_out = nn.GroupNorm(num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6)
|
||||
self.conv_norm_out = nn.GroupNorm(
|
||||
num_channels=block_out_channels[0], num_groups=norm_num_groups, eps=1e-6
|
||||
)
|
||||
self.conv_act = nn.SiLU()
|
||||
self.conv_out = CausalConv3d(block_out_channels[0], out_channels, kernel_size=3)
|
||||
|
||||
@@ -229,7 +254,6 @@ class DecoderCausal3D(nn.Module):
|
||||
if self.training and self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module):
|
||||
|
||||
def custom_forward(*inputs):
|
||||
return module(*inputs)
|
||||
|
||||
@@ -255,12 +279,16 @@ class DecoderCausal3D(nn.Module):
|
||||
)
|
||||
else:
|
||||
# middle
|
||||
sample = torch.utils.checkpoint.checkpoint(create_custom_forward(self.mid_block), sample, latent_embeds)
|
||||
sample = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(self.mid_block), sample, latent_embeds
|
||||
)
|
||||
sample = sample.to(upscale_dtype)
|
||||
|
||||
# up
|
||||
for up_block in self.up_blocks:
|
||||
sample = torch.utils.checkpoint.checkpoint(create_custom_forward(up_block), sample, latent_embeds)
|
||||
sample = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(up_block), sample, latent_embeds
|
||||
)
|
||||
else:
|
||||
# middle
|
||||
sample = self.mid_block(sample, latent_embeds)
|
||||
@@ -282,7 +310,6 @@ class DecoderCausal3D(nn.Module):
|
||||
|
||||
|
||||
class DiagonalGaussianDistribution(object):
|
||||
|
||||
def __init__(self, parameters: torch.Tensor, deterministic: bool = False):
|
||||
if parameters.ndim == 3:
|
||||
dim = 2 # (B, L, C)
|
||||
@@ -297,9 +324,9 @@ class DiagonalGaussianDistribution(object):
|
||||
self.std = torch.exp(0.5 * self.logvar)
|
||||
self.var = torch.exp(self.logvar)
|
||||
if self.deterministic:
|
||||
self.var = self.std = torch.zeros_like(self.mean,
|
||||
device=self.parameters.device,
|
||||
dtype=self.parameters.dtype)
|
||||
self.var = self.std = torch.zeros_like(
|
||||
self.mean, device=self.parameters.device, dtype=self.parameters.dtype
|
||||
)
|
||||
|
||||
def sample(self, generator: Optional[torch.Generator] = None) -> torch.FloatTensor:
|
||||
# make sure sample is on the same device as the parameters and has same dtype
|
||||
@@ -324,12 +351,17 @@ class DiagonalGaussianDistribution(object):
|
||||
)
|
||||
else:
|
||||
return 0.5 * torch.sum(
|
||||
torch.pow(self.mean - other.mean, 2) / other.var + self.var / other.var - 1.0 - self.logvar +
|
||||
other.logvar,
|
||||
torch.pow(self.mean - other.mean, 2) / other.var
|
||||
+ self.var / other.var
|
||||
- 1.0
|
||||
- self.logvar
|
||||
+ other.logvar,
|
||||
dim=reduce_dim,
|
||||
)
|
||||
|
||||
def nll(self, sample: torch.Tensor, dims: Tuple[int, ...] = [1, 2, 3]) -> torch.Tensor:
|
||||
def nll(
|
||||
self, sample: torch.Tensor, dims: Tuple[int, ...] = [1, 2, 3]
|
||||
) -> torch.Tensor:
|
||||
if self.deterministic:
|
||||
return torch.Tensor([0.0])
|
||||
logtwopi = np.log(2.0 * np.pi)
|
||||
|
||||
@@ -1,836 +0,0 @@
|
||||
# Copyright 2024 The Hunyuan Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
|
||||
from diffusers.models.attention import FeedForward
|
||||
from diffusers.models.attention_processor import Attention, AttentionProcessor
|
||||
from diffusers.models.embeddings import (CombinedTimestepGuidanceTextProjEmbeddings, CombinedTimestepTextProjEmbeddings,
|
||||
get_1d_rotary_pos_embed)
|
||||
from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.models.normalization import AdaLayerNormContinuous, AdaLayerNormZero, AdaLayerNormZeroSingle
|
||||
from diffusers.utils import USE_PEFT_BACKEND, is_torch_version, logging, scale_lora_layers, unscale_lora_layers
|
||||
|
||||
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
|
||||
from fastvideo.utils.communications import all_gather, all_to_all_4D
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
def shrink_head(encoder_state, dim):
|
||||
local_heads = encoder_state.shape[dim] // nccl_info.sp_size
|
||||
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
|
||||
|
||||
|
||||
class HunyuanVideoAttnProcessor2_0:
|
||||
|
||||
def __init__(self):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError(
|
||||
"HunyuanVideoAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
|
||||
sequence_length = hidden_states.size(1)
|
||||
encoder_sequence_length = encoder_hidden_states.size(1)
|
||||
if attn.add_q_proj is None and encoder_hidden_states is not None:
|
||||
hidden_states = torch.cat([hidden_states, encoder_hidden_states], dim=1)
|
||||
|
||||
# 1. QKV projections
|
||||
query = attn.to_q(hidden_states)
|
||||
key = attn.to_k(hidden_states)
|
||||
value = attn.to_v(hidden_states)
|
||||
|
||||
query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
|
||||
# 2. QK normalization
|
||||
if attn.norm_q is not None:
|
||||
query = attn.norm_q(query).to(value)
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key).to(value)
|
||||
|
||||
image_rotary_emb = (
|
||||
shrink_head(image_rotary_emb[0], dim=0),
|
||||
shrink_head(image_rotary_emb[1], dim=0),
|
||||
)
|
||||
|
||||
# 3. Rotational positional embeddings applied to latent stream
|
||||
if image_rotary_emb is not None:
|
||||
from diffusers.models.embeddings import apply_rotary_emb
|
||||
|
||||
if attn.add_q_proj is None and encoder_hidden_states is not None:
|
||||
query = torch.cat(
|
||||
[
|
||||
apply_rotary_emb(query[:, :, :-encoder_hidden_states.shape[1]], image_rotary_emb),
|
||||
query[:, :, -encoder_hidden_states.shape[1]:],
|
||||
],
|
||||
dim=2,
|
||||
)
|
||||
key = torch.cat(
|
||||
[
|
||||
apply_rotary_emb(key[:, :, :-encoder_hidden_states.shape[1]], image_rotary_emb),
|
||||
key[:, :, -encoder_hidden_states.shape[1]:],
|
||||
],
|
||||
dim=2,
|
||||
)
|
||||
else:
|
||||
query = apply_rotary_emb(query, image_rotary_emb)
|
||||
key = apply_rotary_emb(key, image_rotary_emb)
|
||||
|
||||
# 4. Encoder condition QKV projection and normalization
|
||||
if attn.add_q_proj is not None and encoder_hidden_states is not None:
|
||||
encoder_query = attn.add_q_proj(encoder_hidden_states)
|
||||
encoder_key = attn.add_k_proj(encoder_hidden_states)
|
||||
encoder_value = attn.add_v_proj(encoder_hidden_states)
|
||||
|
||||
encoder_query = encoder_query.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
encoder_key = encoder_key.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
encoder_value = encoder_value.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
|
||||
if attn.norm_added_q is not None:
|
||||
encoder_query = attn.norm_added_q(encoder_query).to(encoder_value)
|
||||
if attn.norm_added_k is not None:
|
||||
encoder_key = attn.norm_added_k(encoder_key).to(encoder_value)
|
||||
|
||||
query = torch.cat([query, encoder_query], dim=2)
|
||||
key = torch.cat([key, encoder_key], dim=2)
|
||||
value = torch.cat([value, encoder_value], dim=2)
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
query_img, query_txt = query[:, :, :sequence_length, :], query[:, :, sequence_length:, :]
|
||||
key_img, key_txt = key[:, :, :sequence_length, :], key[:, :, sequence_length:, :]
|
||||
value_img, value_txt = value[:, :, :sequence_length, :], value[:, :, sequence_length:, :]
|
||||
query_img = all_to_all_4D(query_img, scatter_dim=1, gather_dim=2) #
|
||||
key_img = all_to_all_4D(key_img, scatter_dim=1, gather_dim=2)
|
||||
value_img = all_to_all_4D(value_img, scatter_dim=1, gather_dim=2)
|
||||
|
||||
query_txt = shrink_head(query_txt, dim=1)
|
||||
key_txt = shrink_head(key_txt, dim=1)
|
||||
value_txt = shrink_head(value_txt, dim=1)
|
||||
query = torch.cat([query_img, query_txt], dim=2)
|
||||
key = torch.cat([key_img, key_txt], dim=2)
|
||||
value = torch.cat([value_img, value_txt], dim=2)
|
||||
|
||||
query = query.unsqueeze(2)
|
||||
key = key.unsqueeze(2)
|
||||
value = value.unsqueeze(2)
|
||||
qkv = torch.cat([query, key, value], dim=2)
|
||||
qkv = qkv.transpose(1, 3)
|
||||
|
||||
# 5. Attention
|
||||
attention_mask = attention_mask[:, 0, :]
|
||||
seq_len = qkv.shape[1]
|
||||
attn_len = attention_mask.shape[1]
|
||||
attention_mask = F.pad(attention_mask, (seq_len - attn_len, 0), value=True)
|
||||
|
||||
hidden_states = flash_attn_no_pad(qkv, attention_mask, causal=False, dropout_p=0.0, softmax_scale=None)
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
|
||||
(sequence_length * nccl_info.sp_size, encoder_sequence_length), dim=1)
|
||||
hidden_states = all_to_all_4D(hidden_states, scatter_dim=1, gather_dim=2)
|
||||
encoder_hidden_states = all_gather(encoder_hidden_states, dim=2).contiguous()
|
||||
hidden_states = hidden_states.flatten(2, 3)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
encoder_hidden_states = encoder_hidden_states.flatten(2, 3)
|
||||
encoder_hidden_states = encoder_hidden_states.to(query.dtype)
|
||||
else:
|
||||
hidden_states = hidden_states.flatten(2, 3)
|
||||
hidden_states = hidden_states.to(query.dtype)
|
||||
|
||||
# 6. Output projection
|
||||
if encoder_hidden_states is not None:
|
||||
hidden_states, encoder_hidden_states = (
|
||||
hidden_states[:, :-encoder_hidden_states.shape[1]],
|
||||
hidden_states[:, -encoder_hidden_states.shape[1]:],
|
||||
)
|
||||
|
||||
if encoder_hidden_states is not None:
|
||||
if getattr(attn, "to_out", None) is not None:
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
|
||||
if getattr(attn, "to_add_out", None) is not None:
|
||||
encoder_hidden_states = attn.to_add_out(encoder_hidden_states)
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoPatchEmbed(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: Union[int, Tuple[int, int, int]] = 16,
|
||||
in_chans: int = 3,
|
||||
embed_dim: int = 768,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
patch_size = (patch_size, patch_size, patch_size) if isinstance(patch_size, int) else patch_size
|
||||
self.proj = nn.Conv3d(in_chans, embed_dim, kernel_size=patch_size, stride=patch_size)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.proj(hidden_states)
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2) # BCFHW -> BNC
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoAdaNorm(nn.Module):
|
||||
|
||||
def __init__(self, in_features: int, out_features: Optional[int] = None) -> None:
|
||||
super().__init__()
|
||||
|
||||
out_features = out_features or 2 * in_features
|
||||
self.linear = nn.Linear(in_features, out_features)
|
||||
self.nonlinearity = nn.SiLU()
|
||||
|
||||
def forward(self,
|
||||
temb: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
temb = self.linear(self.nonlinearity(temb))
|
||||
gate_msa, gate_mlp = temb.chunk(2, dim=1)
|
||||
gate_msa, gate_mlp = gate_msa.unsqueeze(1), gate_mlp.unsqueeze(1)
|
||||
return gate_msa, gate_mlp
|
||||
|
||||
|
||||
class HunyuanVideoIndividualTokenRefinerBlock(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
mlp_width_ratio: str = 4.0,
|
||||
mlp_drop_rate: float = 0.0,
|
||||
attention_bias: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
hidden_size = num_attention_heads * attention_head_dim
|
||||
|
||||
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6)
|
||||
self.attn = Attention(
|
||||
query_dim=hidden_size,
|
||||
cross_attention_dim=None,
|
||||
heads=num_attention_heads,
|
||||
dim_head=attention_head_dim,
|
||||
bias=attention_bias,
|
||||
)
|
||||
|
||||
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6)
|
||||
self.ff = FeedForward(hidden_size, mult=mlp_width_ratio, activation_fn="linear-silu", dropout=mlp_drop_rate)
|
||||
|
||||
self.norm_out = HunyuanVideoAdaNorm(hidden_size, 2 * hidden_size)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
norm_hidden_states = self.norm1(hidden_states)
|
||||
|
||||
attn_output = self.attn(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=None,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
|
||||
gate_msa, gate_mlp = self.norm_out(temb)
|
||||
hidden_states = hidden_states + attn_output * gate_msa
|
||||
|
||||
ff_output = self.ff(self.norm2(hidden_states))
|
||||
hidden_states = hidden_states + ff_output * gate_mlp
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoIndividualTokenRefiner(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
num_layers: int,
|
||||
mlp_width_ratio: float = 4.0,
|
||||
mlp_drop_rate: float = 0.0,
|
||||
attention_bias: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.refiner_blocks = nn.ModuleList([
|
||||
HunyuanVideoIndividualTokenRefinerBlock(
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
mlp_width_ratio=mlp_width_ratio,
|
||||
mlp_drop_rate=mlp_drop_rate,
|
||||
attention_bias=attention_bias,
|
||||
) for _ in range(num_layers)
|
||||
])
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
) -> None:
|
||||
self_attn_mask = None
|
||||
if attention_mask is not None:
|
||||
batch_size = attention_mask.shape[0]
|
||||
seq_len = attention_mask.shape[1]
|
||||
attention_mask = attention_mask.to(hidden_states.device).bool()
|
||||
self_attn_mask_1 = attention_mask.view(batch_size, 1, 1, seq_len).repeat(1, 1, seq_len, 1)
|
||||
self_attn_mask_2 = self_attn_mask_1.transpose(2, 3)
|
||||
self_attn_mask = (self_attn_mask_1 & self_attn_mask_2).bool()
|
||||
self_attn_mask[:, :, :, 0] = True
|
||||
|
||||
for block in self.refiner_blocks:
|
||||
hidden_states = block(hidden_states, temb, self_attn_mask)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoTokenRefiner(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
num_layers: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
mlp_drop_rate: float = 0.0,
|
||||
attention_bias: bool = True,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
hidden_size = num_attention_heads * attention_head_dim
|
||||
|
||||
self.time_text_embed = CombinedTimestepTextProjEmbeddings(embedding_dim=hidden_size,
|
||||
pooled_projection_dim=in_channels)
|
||||
self.proj_in = nn.Linear(in_channels, hidden_size, bias=True)
|
||||
self.token_refiner = HunyuanVideoIndividualTokenRefiner(
|
||||
num_attention_heads=num_attention_heads,
|
||||
attention_head_dim=attention_head_dim,
|
||||
num_layers=num_layers,
|
||||
mlp_width_ratio=mlp_ratio,
|
||||
mlp_drop_rate=mlp_drop_rate,
|
||||
attention_bias=attention_bias,
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
attention_mask: Optional[torch.LongTensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if attention_mask is None:
|
||||
pooled_projections = hidden_states.mean(dim=1)
|
||||
else:
|
||||
original_dtype = hidden_states.dtype
|
||||
mask_float = attention_mask.float().unsqueeze(-1)
|
||||
pooled_projections = (hidden_states * mask_float).sum(dim=1) / mask_float.sum(dim=1)
|
||||
pooled_projections = pooled_projections.to(original_dtype)
|
||||
|
||||
temb = self.time_text_embed(timestep, pooled_projections)
|
||||
hidden_states = self.proj_in(hidden_states)
|
||||
hidden_states = self.token_refiner(hidden_states, temb, attention_mask)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoRotaryPosEmbed(nn.Module):
|
||||
|
||||
def __init__(self, patch_size: int, patch_size_t: int, rope_dim: List[int], theta: float = 256.0) -> None:
|
||||
super().__init__()
|
||||
|
||||
self.patch_size = patch_size
|
||||
self.patch_size_t = patch_size_t
|
||||
self.rope_dim = rope_dim
|
||||
self.theta = theta
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
rope_sizes = [
|
||||
num_frames * nccl_info.sp_size // self.patch_size_t, height // self.patch_size, width // self.patch_size
|
||||
]
|
||||
|
||||
axes_grids = []
|
||||
for i in range(3):
|
||||
# Note: The following line diverges from original behaviour. We create the grid on the device, whereas
|
||||
# original implementation creates it on CPU and then moves it to device. This results in numerical
|
||||
# differences in layerwise debugging outputs, but visually it is the same.
|
||||
grid = torch.arange(0, rope_sizes[i], device=hidden_states.device, dtype=torch.float32)
|
||||
axes_grids.append(grid)
|
||||
grid = torch.meshgrid(*axes_grids, indexing="ij") # [W, H, T]
|
||||
grid = torch.stack(grid, dim=0) # [3, W, H, T]
|
||||
|
||||
freqs = []
|
||||
for i in range(3):
|
||||
freq = get_1d_rotary_pos_embed(self.rope_dim[i], grid[i].reshape(-1), self.theta, use_real=True)
|
||||
freqs.append(freq)
|
||||
|
||||
freqs_cos = torch.cat([f[0] for f in freqs], dim=1) # (W * H * T, D / 2)
|
||||
freqs_sin = torch.cat([f[1] for f in freqs], dim=1) # (W * H * T, D / 2)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
class HunyuanVideoSingleTransformerBlock(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
qk_norm: str = "rms_norm",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
hidden_size = num_attention_heads * attention_head_dim
|
||||
mlp_dim = int(hidden_size * mlp_ratio)
|
||||
|
||||
self.attn = Attention(
|
||||
query_dim=hidden_size,
|
||||
cross_attention_dim=None,
|
||||
dim_head=attention_head_dim,
|
||||
heads=num_attention_heads,
|
||||
out_dim=hidden_size,
|
||||
bias=True,
|
||||
processor=HunyuanVideoAttnProcessor2_0(),
|
||||
qk_norm=qk_norm,
|
||||
eps=1e-6,
|
||||
pre_only=True,
|
||||
)
|
||||
|
||||
self.norm = AdaLayerNormZeroSingle(hidden_size, norm_type="layer_norm")
|
||||
self.proj_mlp = nn.Linear(hidden_size, mlp_dim)
|
||||
self.act_mlp = nn.GELU(approximate="tanh")
|
||||
self.proj_out = nn.Linear(hidden_size + mlp_dim, hidden_size)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
image_rotary_emb: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
) -> torch.Tensor:
|
||||
text_seq_length = encoder_hidden_states.shape[1]
|
||||
hidden_states = torch.cat([hidden_states, encoder_hidden_states], dim=1)
|
||||
|
||||
residual = hidden_states
|
||||
|
||||
# 1. Input normalization
|
||||
norm_hidden_states, gate = self.norm(hidden_states, emb=temb)
|
||||
mlp_hidden_states = self.act_mlp(self.proj_mlp(norm_hidden_states))
|
||||
|
||||
norm_hidden_states, norm_encoder_hidden_states = (
|
||||
norm_hidden_states[:, :-text_seq_length, :],
|
||||
norm_hidden_states[:, -text_seq_length:, :],
|
||||
)
|
||||
|
||||
# 2. Attention
|
||||
attn_output, context_attn_output = self.attn(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
image_rotary_emb=image_rotary_emb,
|
||||
)
|
||||
attn_output = torch.cat([attn_output, context_attn_output], dim=1)
|
||||
|
||||
# 3. Modulation and residual connection
|
||||
hidden_states = torch.cat([attn_output, mlp_hidden_states], dim=2)
|
||||
hidden_states = gate.unsqueeze(1) * self.proj_out(hidden_states)
|
||||
hidden_states = hidden_states + residual
|
||||
|
||||
hidden_states, encoder_hidden_states = (
|
||||
hidden_states[:, :-text_seq_length, :],
|
||||
hidden_states[:, -text_seq_length:, :],
|
||||
)
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoTransformerBlock(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_attention_heads: int,
|
||||
attention_head_dim: int,
|
||||
mlp_ratio: float,
|
||||
qk_norm: str = "rms_norm",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
hidden_size = num_attention_heads * attention_head_dim
|
||||
|
||||
self.norm1 = AdaLayerNormZero(hidden_size, norm_type="layer_norm")
|
||||
self.norm1_context = AdaLayerNormZero(hidden_size, norm_type="layer_norm")
|
||||
|
||||
self.attn = Attention(
|
||||
query_dim=hidden_size,
|
||||
cross_attention_dim=None,
|
||||
added_kv_proj_dim=hidden_size,
|
||||
dim_head=attention_head_dim,
|
||||
heads=num_attention_heads,
|
||||
out_dim=hidden_size,
|
||||
context_pre_only=False,
|
||||
bias=True,
|
||||
processor=HunyuanVideoAttnProcessor2_0(),
|
||||
qk_norm=qk_norm,
|
||||
eps=1e-6,
|
||||
)
|
||||
|
||||
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.ff = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu-approximate")
|
||||
|
||||
self.norm2_context = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6)
|
||||
self.ff_context = FeedForward(hidden_size, mult=mlp_ratio, activation_fn="gelu-approximate")
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
freqs_cis: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# 1. Input normalization
|
||||
norm_hidden_states, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.norm1(hidden_states, emb=temb)
|
||||
norm_encoder_hidden_states, c_gate_msa, c_shift_mlp, c_scale_mlp, c_gate_mlp = self.norm1_context(
|
||||
encoder_hidden_states, emb=temb)
|
||||
|
||||
# 2. Joint attention
|
||||
attn_output, context_attn_output = self.attn(
|
||||
hidden_states=norm_hidden_states,
|
||||
encoder_hidden_states=norm_encoder_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
image_rotary_emb=freqs_cis,
|
||||
)
|
||||
|
||||
# 3. Modulation and residual connection
|
||||
hidden_states = hidden_states + attn_output * gate_msa.unsqueeze(1)
|
||||
encoder_hidden_states = encoder_hidden_states + context_attn_output * c_gate_msa.unsqueeze(1)
|
||||
|
||||
norm_hidden_states = self.norm2(hidden_states)
|
||||
norm_encoder_hidden_states = self.norm2_context(encoder_hidden_states)
|
||||
|
||||
norm_hidden_states = norm_hidden_states * (1 + scale_mlp[:, None]) + shift_mlp[:, None]
|
||||
norm_encoder_hidden_states = norm_encoder_hidden_states * (1 + c_scale_mlp[:, None]) + c_shift_mlp[:, None]
|
||||
|
||||
# 4. Feed-forward
|
||||
ff_output = self.ff(norm_hidden_states)
|
||||
context_ff_output = self.ff_context(norm_encoder_hidden_states)
|
||||
|
||||
hidden_states = hidden_states + gate_mlp.unsqueeze(1) * ff_output
|
||||
encoder_hidden_states = encoder_hidden_states + c_gate_mlp.unsqueeze(1) * context_ff_output
|
||||
|
||||
return hidden_states, encoder_hidden_states
|
||||
|
||||
|
||||
class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin):
|
||||
r"""
|
||||
A Transformer model for video-like data used in [HunyuanVideo](https://huggingface.co/tencent/HunyuanVideo).
|
||||
|
||||
Args:
|
||||
in_channels (`int`, defaults to `16`):
|
||||
The number of channels in the input.
|
||||
out_channels (`int`, defaults to `16`):
|
||||
The number of channels in the output.
|
||||
num_attention_heads (`int`, defaults to `24`):
|
||||
The number of heads to use for multi-head attention.
|
||||
attention_head_dim (`int`, defaults to `128`):
|
||||
The number of channels in each head.
|
||||
num_layers (`int`, defaults to `20`):
|
||||
The number of layers of dual-stream blocks to use.
|
||||
num_single_layers (`int`, defaults to `40`):
|
||||
The number of layers of single-stream blocks to use.
|
||||
num_refiner_layers (`int`, defaults to `2`):
|
||||
The number of layers of refiner blocks to use.
|
||||
mlp_ratio (`float`, defaults to `4.0`):
|
||||
The ratio of the hidden layer size to the input size in the feedforward network.
|
||||
patch_size (`int`, defaults to `2`):
|
||||
The size of the spatial patches to use in the patch embedding layer.
|
||||
patch_size_t (`int`, defaults to `1`):
|
||||
The size of the tmeporal patches to use in the patch embedding layer.
|
||||
qk_norm (`str`, defaults to `rms_norm`):
|
||||
The normalization to use for the query and key projections in the attention layers.
|
||||
guidance_embeds (`bool`, defaults to `True`):
|
||||
Whether to use guidance embeddings in the model.
|
||||
text_embed_dim (`int`, defaults to `4096`):
|
||||
Input dimension of text embeddings from the text encoder.
|
||||
pooled_projection_dim (`int`, defaults to `768`):
|
||||
The dimension of the pooled projection of the text embeddings.
|
||||
rope_theta (`float`, defaults to `256.0`):
|
||||
The value of theta to use in the RoPE layer.
|
||||
rope_axes_dim (`Tuple[int]`, defaults to `(16, 56, 56)`):
|
||||
The dimensions of the axes to use in the RoPE layer.
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
in_channels: int = 16,
|
||||
out_channels: int = 16,
|
||||
num_attention_heads: int = 24,
|
||||
attention_head_dim: int = 128,
|
||||
num_layers: int = 20,
|
||||
num_single_layers: int = 40,
|
||||
num_refiner_layers: int = 2,
|
||||
mlp_ratio: float = 4.0,
|
||||
patch_size: int = 2,
|
||||
patch_size_t: int = 1,
|
||||
qk_norm: str = "rms_norm",
|
||||
guidance_embeds: bool = True,
|
||||
text_embed_dim: int = 4096,
|
||||
pooled_projection_dim: int = 768,
|
||||
rope_theta: float = 256.0,
|
||||
rope_axes_dim: Tuple[int] = (16, 56, 56),
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
out_channels = out_channels or in_channels
|
||||
|
||||
# 1. Latent and condition embedders
|
||||
self.x_embedder = HunyuanVideoPatchEmbed((patch_size_t, patch_size, patch_size), in_channels, inner_dim)
|
||||
self.context_embedder = HunyuanVideoTokenRefiner(text_embed_dim,
|
||||
num_attention_heads,
|
||||
attention_head_dim,
|
||||
num_layers=num_refiner_layers)
|
||||
self.time_text_embed = CombinedTimestepGuidanceTextProjEmbeddings(inner_dim, pooled_projection_dim)
|
||||
|
||||
# 2. RoPE
|
||||
self.rope = HunyuanVideoRotaryPosEmbed(patch_size, patch_size_t, rope_axes_dim, rope_theta)
|
||||
|
||||
# 3. Dual stream transformer blocks
|
||||
self.transformer_blocks = nn.ModuleList([
|
||||
HunyuanVideoTransformerBlock(num_attention_heads, attention_head_dim, mlp_ratio=mlp_ratio, qk_norm=qk_norm)
|
||||
for _ in range(num_layers)
|
||||
])
|
||||
|
||||
# 4. Single stream transformer blocks
|
||||
self.single_transformer_blocks = nn.ModuleList([
|
||||
HunyuanVideoSingleTransformerBlock(num_attention_heads,
|
||||
attention_head_dim,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qk_norm=qk_norm) for _ in range(num_single_layers)
|
||||
])
|
||||
|
||||
# 5. Output projection
|
||||
self.norm_out = AdaLayerNormContinuous(inner_dim, inner_dim, elementwise_affine=False, eps=1e-6)
|
||||
self.proj_out = nn.Linear(inner_dim, patch_size_t * patch_size * patch_size * out_channels)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
@property
|
||||
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.attn_processors
|
||||
def attn_processors(self) -> Dict[str, AttentionProcessor]:
|
||||
r"""
|
||||
Returns:
|
||||
`dict` of attention processors: A dictionary containing all attention processors used in the model with
|
||||
indexed by its weight name.
|
||||
"""
|
||||
# set recursively
|
||||
processors = {}
|
||||
|
||||
def fn_recursive_add_processors(name: str, module: torch.nn.Module, processors: Dict[str, AttentionProcessor]):
|
||||
if hasattr(module, "get_processor"):
|
||||
processors[f"{name}.processor"] = module.get_processor()
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
|
||||
|
||||
return processors
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_add_processors(name, module, processors)
|
||||
|
||||
return processors
|
||||
|
||||
# Copied from diffusers.models.unets.unet_2d_condition.UNet2DConditionModel.set_attn_processor
|
||||
def set_attn_processor(self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]]):
|
||||
r"""
|
||||
Sets the attention processor to use to compute attention.
|
||||
|
||||
Parameters:
|
||||
processor (`dict` of `AttentionProcessor` or only `AttentionProcessor`):
|
||||
The instantiated processor class or a dictionary of processor classes that will be set as the processor
|
||||
for **all** `Attention` layers.
|
||||
|
||||
If `processor` is a dict, the key needs to define the path to the corresponding cross attention
|
||||
processor. This is strongly recommended when setting trainable attention processors.
|
||||
|
||||
"""
|
||||
count = len(self.attn_processors.keys())
|
||||
|
||||
if isinstance(processor, dict) and len(processor) != count:
|
||||
raise ValueError(
|
||||
f"A dict of processors was passed, but the number of processors {len(processor)} does not match the"
|
||||
f" number of attention layers: {count}. Please make sure to pass {count} processor classes.")
|
||||
|
||||
def fn_recursive_attn_processor(name: str, module: torch.nn.Module, processor):
|
||||
if hasattr(module, "set_processor"):
|
||||
if not isinstance(processor, dict):
|
||||
module.set_processor(processor)
|
||||
else:
|
||||
module.set_processor(processor.pop(f"{name}.processor"))
|
||||
|
||||
for sub_name, child in module.named_children():
|
||||
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_attn_processor(name, module, processor)
|
||||
|
||||
def _set_gradient_checkpointing(self, module, value=False):
|
||||
if hasattr(module, "gradient_checkpointing"):
|
||||
module.gradient_checkpointing = value
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_attention_mask: torch.Tensor,
|
||||
guidance: torch.Tensor = None,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
return_dict: bool = True,
|
||||
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||
if guidance is None:
|
||||
guidance = torch.tensor([6016.0], device=hidden_states.device, dtype=torch.bfloat16)
|
||||
|
||||
if attention_kwargs is not None:
|
||||
attention_kwargs = attention_kwargs.copy()
|
||||
lora_scale = attention_kwargs.pop("scale", 1.0)
|
||||
else:
|
||||
lora_scale = 1.0
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# weight the lora layers by setting `lora_scale` for each PEFT layer
|
||||
scale_lora_layers(self, lora_scale)
|
||||
else:
|
||||
if attention_kwargs is not None and attention_kwargs.get("scale", None) is not None:
|
||||
logger.warning("Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective.")
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p, p_t = self.config.patch_size, self.config.patch_size_t
|
||||
post_patch_num_frames = num_frames // p_t
|
||||
post_patch_height = height // p
|
||||
post_patch_width = width // p
|
||||
|
||||
pooled_projections = encoder_hidden_states[:, 0, :self.config.pooled_projection_dim]
|
||||
encoder_hidden_states = encoder_hidden_states[:, 1:]
|
||||
|
||||
# 1. RoPE
|
||||
image_rotary_emb = self.rope(hidden_states)
|
||||
|
||||
# 2. Conditional embeddings
|
||||
temb = self.time_text_embed(timestep, guidance, pooled_projections)
|
||||
hidden_states = self.x_embedder(hidden_states)
|
||||
encoder_hidden_states = self.context_embedder(encoder_hidden_states, timestep, encoder_attention_mask)
|
||||
|
||||
# 3. Attention mask preparation
|
||||
latent_sequence_length = hidden_states.shape[1]
|
||||
condition_sequence_length = encoder_hidden_states.shape[1]
|
||||
sequence_length = latent_sequence_length + condition_sequence_length
|
||||
attention_mask = torch.zeros(batch_size,
|
||||
sequence_length,
|
||||
sequence_length,
|
||||
device=hidden_states.device,
|
||||
dtype=torch.bool) # [B, N, N]
|
||||
|
||||
effective_condition_sequence_length = encoder_attention_mask.sum(dim=1, dtype=torch.int)
|
||||
effective_sequence_length = latent_sequence_length + effective_condition_sequence_length
|
||||
|
||||
for i in range(batch_size):
|
||||
attention_mask[i, :effective_sequence_length[i], :effective_sequence_length[i]] = True
|
||||
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
|
||||
def create_custom_forward(module, return_dict=None):
|
||||
|
||||
def custom_forward(*inputs):
|
||||
if return_dict is not None:
|
||||
return module(*inputs, return_dict=return_dict)
|
||||
else:
|
||||
return module(*inputs)
|
||||
|
||||
return custom_forward
|
||||
|
||||
ckpt_kwargs: Dict[str, Any] = {"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {}
|
||||
|
||||
for block in self.transformer_blocks:
|
||||
hidden_states, encoder_hidden_states = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(block),
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
temb,
|
||||
attention_mask,
|
||||
image_rotary_emb,
|
||||
**ckpt_kwargs,
|
||||
)
|
||||
|
||||
for block in self.single_transformer_blocks:
|
||||
hidden_states, encoder_hidden_states = torch.utils.checkpoint.checkpoint(
|
||||
create_custom_forward(block),
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
temb,
|
||||
attention_mask,
|
||||
image_rotary_emb,
|
||||
**ckpt_kwargs,
|
||||
)
|
||||
|
||||
else:
|
||||
for block in self.transformer_blocks:
|
||||
hidden_states, encoder_hidden_states = block(hidden_states, encoder_hidden_states, temb, attention_mask,
|
||||
image_rotary_emb)
|
||||
|
||||
for block in self.single_transformer_blocks:
|
||||
hidden_states, encoder_hidden_states = block(hidden_states, encoder_hidden_states, temb, attention_mask,
|
||||
image_rotary_emb)
|
||||
|
||||
# 5. Output projection
|
||||
hidden_states = self.norm_out(hidden_states, temb)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames, post_patch_height, post_patch_width,
|
||||
-1, p_t, p, p)
|
||||
hidden_states = hidden_states.permute(0, 4, 1, 5, 2, 6, 3, 7)
|
||||
hidden_states = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# remove `lora_scale` from each PEFT layer
|
||||
unscale_lora_layers(self, lora_scale)
|
||||
|
||||
if not return_dict:
|
||||
return (hidden_states, )
|
||||
|
||||
return Transformer2DModelOutput(sample=hidden_states)
|
||||
@@ -1,691 +0,0 @@
|
||||
# Copyright 2024 The HunyuanVideo Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import inspect
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
|
||||
from diffusers.loaders import HunyuanVideoLoraLoaderMixin
|
||||
from diffusers.models import AutoencoderKLHunyuanVideo, HunyuanVideoTransformer3DModel
|
||||
from diffusers.pipelines.hunyuan_video.pipeline_output import HunyuanVideoPipelineOutput
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import logging, replace_example_docstring
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from einops import rearrange
|
||||
from transformers import CLIPTextModel, CLIPTokenizer, LlamaModel, LlamaTokenizerFast
|
||||
|
||||
from fastvideo.utils.communications import all_gather
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
EXAMPLE_DOC_STRING = """
|
||||
Examples:
|
||||
```python
|
||||
>>> import torch
|
||||
>>> from diffusers import HunyuanVideoPipeline, HunyuanVideoTransformer3DModel
|
||||
>>> from diffusers.utils import export_to_video
|
||||
|
||||
>>> model_id = "tencent/HunyuanVideo"
|
||||
>>> transformer = HunyuanVideoTransformer3DModel.from_pretrained(
|
||||
... model_id, subfolder="transformer", torch_dtype=torch.bfloat16
|
||||
... )
|
||||
>>> pipe = HunyuanVideoPipeline.from_pretrained(model_id, transformer=transformer, torch_dtype=torch.float16)
|
||||
>>> pipe.vae.enable_tiling()
|
||||
>>> pipe.to("cuda")
|
||||
|
||||
>>> output = pipe(
|
||||
... prompt="A cat walks on the grass, realistic",
|
||||
... height=320,
|
||||
... width=512,
|
||||
... num_frames=61,
|
||||
... num_inference_steps=30,
|
||||
... ).frames[0]
|
||||
>>> export_to_video(output, "output.mp4", fps=15)
|
||||
```
|
||||
"""
|
||||
|
||||
DEFAULT_PROMPT_TEMPLATE = {
|
||||
"template": ("<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
|
||||
"1. The main content and theme of the video."
|
||||
"2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects."
|
||||
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
|
||||
"4. background environment, light, style and atmosphere."
|
||||
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
|
||||
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"),
|
||||
"crop_start":
|
||||
95,
|
||||
}
|
||||
|
||||
|
||||
# Copied from diffusers.pipelines.stable_diffusion.pipeline_stable_diffusion.retrieve_timesteps
|
||||
def retrieve_timesteps(
|
||||
scheduler,
|
||||
num_inference_steps: Optional[int] = None,
|
||||
device: Optional[Union[str, torch.device]] = None,
|
||||
timesteps: Optional[List[int]] = None,
|
||||
sigmas: Optional[List[float]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
r"""
|
||||
Calls the scheduler's `set_timesteps` method and retrieves timesteps from the scheduler after the call. Handles
|
||||
custom timesteps. Any kwargs will be supplied to `scheduler.set_timesteps`.
|
||||
|
||||
Args:
|
||||
scheduler (`SchedulerMixin`):
|
||||
The scheduler to get timesteps from.
|
||||
num_inference_steps (`int`):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model. If used, `timesteps`
|
||||
must be `None`.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
timesteps (`List[int]`, *optional*):
|
||||
Custom timesteps used to override the timestep spacing strategy of the scheduler. If `timesteps` is passed,
|
||||
`num_inference_steps` and `sigmas` must be `None`.
|
||||
sigmas (`List[float]`, *optional*):
|
||||
Custom sigmas used to override the timestep spacing strategy of the scheduler. If `sigmas` is passed,
|
||||
`num_inference_steps` and `timesteps` must be `None`.
|
||||
|
||||
Returns:
|
||||
`Tuple[torch.Tensor, int]`: A tuple where the first element is the timestep schedule from the scheduler and the
|
||||
second element is the number of inference steps.
|
||||
"""
|
||||
if timesteps is not None and sigmas is not None:
|
||||
raise ValueError("Only one of `timesteps` or `sigmas` can be passed. Please choose one to set custom values")
|
||||
if timesteps is not None:
|
||||
accepts_timesteps = "timesteps" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
|
||||
if not accepts_timesteps:
|
||||
raise ValueError(
|
||||
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
|
||||
f" timestep schedules. Please check whether you are using the correct scheduler.")
|
||||
scheduler.set_timesteps(timesteps=timesteps, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
elif sigmas is not None:
|
||||
accept_sigmas = "sigmas" in set(inspect.signature(scheduler.set_timesteps).parameters.keys())
|
||||
if not accept_sigmas:
|
||||
raise ValueError(
|
||||
f"The current scheduler class {scheduler.__class__}'s `set_timesteps` does not support custom"
|
||||
f" sigmas schedules. Please check whether you are using the correct scheduler.")
|
||||
scheduler.set_timesteps(sigmas=sigmas, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
num_inference_steps = len(timesteps)
|
||||
else:
|
||||
scheduler.set_timesteps(num_inference_steps, device=device, **kwargs)
|
||||
timesteps = scheduler.timesteps
|
||||
return timesteps, num_inference_steps
|
||||
|
||||
|
||||
class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
r"""
|
||||
Pipeline for text-to-video generation using HunyuanVideo.
|
||||
|
||||
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods
|
||||
implemented for all pipelines (downloading, saving, running on a particular device, etc.).
|
||||
|
||||
Args:
|
||||
text_encoder ([`LlamaModel`]):
|
||||
[Llava Llama3-8B](https://huggingface.co/xtuner/llava-llama-3-8b-v1_1-transformers).
|
||||
tokenizer_2 (`LlamaTokenizer`):
|
||||
Tokenizer from [Llava Llama3-8B](https://huggingface.co/xtuner/llava-llama-3-8b-v1_1-transformers).
|
||||
transformer ([`HunyuanVideoTransformer3DModel`]):
|
||||
Conditional Transformer to denoise the encoded image latents.
|
||||
scheduler ([`FlowMatchEulerDiscreteScheduler`]):
|
||||
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
|
||||
vae ([`AutoencoderKLHunyuanVideo`]):
|
||||
Variational Auto-Encoder (VAE) Model to encode and decode videos to and from latent representations.
|
||||
text_encoder_2 ([`CLIPTextModel`]):
|
||||
[CLIP](https://huggingface.co/docs/transformers/model_doc/clip#transformers.CLIPTextModel), specifically
|
||||
the [clip-vit-large-patch14](https://huggingface.co/openai/clip-vit-large-patch14) variant.
|
||||
tokenizer_2 (`CLIPTokenizer`):
|
||||
Tokenizer of class
|
||||
[CLIPTokenizer](https://huggingface.co/docs/transformers/en/model_doc/clip#transformers.CLIPTokenizer).
|
||||
"""
|
||||
|
||||
model_cpu_offload_seq = "text_encoder->text_encoder_2->transformer->vae"
|
||||
_callback_tensor_inputs = ["latents", "prompt_embeds"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
text_encoder: LlamaModel,
|
||||
tokenizer: LlamaTokenizerFast,
|
||||
transformer: HunyuanVideoTransformer3DModel,
|
||||
vae: AutoencoderKLHunyuanVideo,
|
||||
scheduler: FlowMatchEulerDiscreteScheduler,
|
||||
text_encoder_2: CLIPTextModel,
|
||||
tokenizer_2: CLIPTokenizer,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
transformer=transformer,
|
||||
scheduler=scheduler,
|
||||
text_encoder_2=text_encoder_2,
|
||||
tokenizer_2=tokenizer_2,
|
||||
)
|
||||
|
||||
self.vae_scale_factor_temporal = (self.vae.temporal_compression_ratio
|
||||
if hasattr(self, "vae") and self.vae is not None else 4)
|
||||
self.vae_scale_factor_spatial = (self.vae.spatial_compression_ratio
|
||||
if hasattr(self, "vae") and self.vae is not None else 8)
|
||||
self.video_processor = VideoProcessor(vae_scale_factor=self.vae_scale_factor_spatial)
|
||||
|
||||
def _get_llama_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
prompt_template: Dict[str, Any],
|
||||
num_videos_per_prompt: int = 1,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
max_sequence_length: int = 256,
|
||||
num_hidden_layers_to_skip: int = 2,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
device = device or self._execution_device
|
||||
dtype = dtype or self.text_encoder.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
prompt = [prompt_template["template"].format(p) for p in prompt]
|
||||
|
||||
crop_start = prompt_template.get("crop_start", None)
|
||||
if crop_start is None:
|
||||
prompt_template_input = self.tokenizer(
|
||||
prompt_template["template"],
|
||||
padding="max_length",
|
||||
return_tensors="pt",
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
return_attention_mask=False,
|
||||
)
|
||||
crop_start = prompt_template_input["input_ids"].shape[-1]
|
||||
# Remove <|eot_id|> token and placeholder {}
|
||||
crop_start -= 2
|
||||
|
||||
max_sequence_length += crop_start
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
max_length=max_sequence_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
return_attention_mask=True,
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids.to(device=device)
|
||||
prompt_attention_mask = text_inputs.attention_mask.to(device=device)
|
||||
|
||||
prompt_embeds = self.text_encoder(
|
||||
input_ids=text_input_ids,
|
||||
attention_mask=prompt_attention_mask,
|
||||
output_hidden_states=True,
|
||||
).hidden_states[-(num_hidden_layers_to_skip + 1)]
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype)
|
||||
|
||||
if crop_start is not None and crop_start > 0:
|
||||
prompt_embeds = prompt_embeds[:, crop_start:]
|
||||
prompt_attention_mask = prompt_attention_mask[:, crop_start:]
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
_, seq_len, _ = prompt_embeds.shape
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
|
||||
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1)
|
||||
prompt_attention_mask = prompt_attention_mask.repeat(1, num_videos_per_prompt)
|
||||
prompt_attention_mask = prompt_attention_mask.view(batch_size * num_videos_per_prompt, seq_len)
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
def _get_clip_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
num_videos_per_prompt: int = 1,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
max_sequence_length: int = 77,
|
||||
) -> torch.Tensor:
|
||||
device = device or self._execution_device
|
||||
dtype = dtype or self.text_encoder_2.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
text_inputs = self.tokenizer_2(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
text_input_ids = text_inputs.input_ids
|
||||
untruncated_ids = self.tokenizer_2(prompt, padding="longest", return_tensors="pt").input_ids
|
||||
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
|
||||
removed_text = self.tokenizer_2.batch_decode(untruncated_ids[:, max_sequence_length - 1:-1])
|
||||
logger.warning("The following part of your input was truncated because CLIP can only handle sequences up to"
|
||||
f" {max_sequence_length} tokens: {removed_text}")
|
||||
|
||||
prompt_embeds = self.text_encoder_2(text_input_ids.to(device), output_hidden_states=False).pooler_output
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt)
|
||||
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, -1)
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
def encode_prompt(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
prompt_2: Union[str, List[str]] = None,
|
||||
prompt_template: Dict[str, Any] = DEFAULT_PROMPT_TEMPLATE,
|
||||
num_videos_per_prompt: int = 1,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
pooled_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
prompt_attention_mask: Optional[torch.Tensor] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
max_sequence_length: int = 256,
|
||||
):
|
||||
|
||||
if prompt_embeds is None:
|
||||
prompt_embeds, prompt_attention_mask = self._get_llama_prompt_embeds(
|
||||
prompt,
|
||||
prompt_template,
|
||||
num_videos_per_prompt,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
max_sequence_length=max_sequence_length,
|
||||
)
|
||||
|
||||
if pooled_prompt_embeds is None:
|
||||
if prompt_2 is None and pooled_prompt_embeds is None:
|
||||
prompt_2 = prompt
|
||||
pooled_prompt_embeds = self._get_clip_prompt_embeds(
|
||||
prompt,
|
||||
num_videos_per_prompt,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
max_sequence_length=77,
|
||||
)
|
||||
|
||||
return prompt_embeds, pooled_prompt_embeds, prompt_attention_mask
|
||||
|
||||
def check_inputs(
|
||||
self,
|
||||
prompt,
|
||||
prompt_2,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds=None,
|
||||
callback_on_step_end_tensor_inputs=None,
|
||||
prompt_template=None,
|
||||
):
|
||||
if height % 16 != 0 or width % 16 != 0:
|
||||
raise ValueError(f"`height` and `width` have to be divisible by 16 but are {height} and {width}.")
|
||||
|
||||
if callback_on_step_end_tensor_inputs is not None and not all(k in self._callback_tensor_inputs
|
||||
for k in callback_on_step_end_tensor_inputs):
|
||||
raise ValueError(
|
||||
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
|
||||
)
|
||||
|
||||
if prompt is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two.")
|
||||
elif prompt_2 is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt_2`: {prompt_2} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two.")
|
||||
elif prompt is None and prompt_embeds is None:
|
||||
raise ValueError(
|
||||
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined.")
|
||||
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
|
||||
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
|
||||
elif prompt_2 is not None and (not isinstance(prompt_2, str) and not isinstance(prompt_2, list)):
|
||||
raise ValueError(f"`prompt_2` has to be of type `str` or `list` but is {type(prompt_2)}")
|
||||
|
||||
if prompt_template is not None:
|
||||
if not isinstance(prompt_template, dict):
|
||||
raise ValueError(f"`prompt_template` has to be of type `dict` but is {type(prompt_template)}")
|
||||
if "template" not in prompt_template:
|
||||
raise ValueError(
|
||||
f"`prompt_template` has to contain a key `template` but only found {prompt_template.keys()}")
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
batch_size: int,
|
||||
num_channels_latents: 32,
|
||||
height: int = 720,
|
||||
width: int = 1280,
|
||||
num_frames: int = 129,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if latents is not None:
|
||||
return latents.to(device=device, dtype=dtype)
|
||||
|
||||
shape = (
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
num_frames,
|
||||
int(height) // self.vae_scale_factor_spatial,
|
||||
int(width) // self.vae_scale_factor_spatial,
|
||||
)
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators.")
|
||||
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
return latents
|
||||
|
||||
def enable_vae_slicing(self):
|
||||
r"""
|
||||
Enable sliced VAE decoding. When this option is enabled, the VAE will split the input tensor in slices to
|
||||
compute decoding in several steps. This is useful to save some memory and allow larger batch sizes.
|
||||
"""
|
||||
self.vae.enable_slicing()
|
||||
|
||||
def disable_vae_slicing(self):
|
||||
r"""
|
||||
Disable sliced VAE decoding. If `enable_vae_slicing` was previously enabled, this method will go back to
|
||||
computing decoding in one step.
|
||||
"""
|
||||
self.vae.disable_slicing()
|
||||
|
||||
def enable_vae_tiling(self):
|
||||
r"""
|
||||
Enable tiled VAE decoding. When this option is enabled, the VAE will split the input tensor into tiles to
|
||||
compute decoding and encoding in several steps. This is useful for saving a large amount of memory and to allow
|
||||
processing larger images.
|
||||
"""
|
||||
self.vae.enable_tiling()
|
||||
|
||||
def disable_vae_tiling(self):
|
||||
r"""
|
||||
Disable tiled VAE decoding. If `enable_vae_tiling` was previously enabled, this method will go back to
|
||||
computing decoding in one step.
|
||||
"""
|
||||
self.vae.disable_tiling()
|
||||
|
||||
@property
|
||||
def guidance_scale(self):
|
||||
return self._guidance_scale
|
||||
|
||||
@property
|
||||
def num_timesteps(self):
|
||||
return self._num_timesteps
|
||||
|
||||
@property
|
||||
def attention_kwargs(self):
|
||||
return self._attention_kwargs
|
||||
|
||||
@property
|
||||
def interrupt(self):
|
||||
return self._interrupt
|
||||
|
||||
@torch.no_grad()
|
||||
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
prompt_2: Union[str, List[str]] = None,
|
||||
height: int = 720,
|
||||
width: int = 1280,
|
||||
num_frames: int = 129,
|
||||
num_inference_steps: int = 50,
|
||||
sigmas: List[float] = None,
|
||||
guidance_scale: float = 6.0,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
pooled_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
prompt_attention_mask: Optional[torch.Tensor] = None,
|
||||
output_type: Optional[str] = "pil",
|
||||
return_dict: bool = True,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
callback_on_step_end: Optional[Union[Callable[[int, int, Dict], None], PipelineCallback,
|
||||
MultiPipelineCallbacks]] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
prompt_template: Dict[str, Any] = DEFAULT_PROMPT_TEMPLATE,
|
||||
max_sequence_length: int = 256,
|
||||
):
|
||||
r"""
|
||||
The call function to the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||
instead.
|
||||
prompt_2 (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to be sent to `tokenizer_2` and `text_encoder_2`. If not defined, `prompt` is
|
||||
will be used instead.
|
||||
height (`int`, defaults to `720`):
|
||||
The height in pixels of the generated image.
|
||||
width (`int`, defaults to `1280`):
|
||||
The width in pixels of the generated image.
|
||||
num_frames (`int`, defaults to `129`):
|
||||
The number of frames in the generated video.
|
||||
num_inference_steps (`int`, defaults to `50`):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
sigmas (`List[float]`, *optional*):
|
||||
Custom sigmas to use for the denoising process with schedulers which support a `sigmas` argument in
|
||||
their `set_timesteps` method. If not defined, the default behavior when `num_inference_steps` is passed
|
||||
will be used.
|
||||
guidance_scale (`float`, defaults to `6.0`):
|
||||
Guidance scale as defined in [Classifier-Free Diffusion Guidance](https://arxiv.org/abs/2207.12598).
|
||||
`guidance_scale` is defined as `w` of equation 2. of [Imagen
|
||||
Paper](https://arxiv.org/pdf/2205.11487.pdf). Guidance scale is enabled by setting `guidance_scale >
|
||||
1`. Higher guidance scale encourages to generate images that are closely linked to the text `prompt`,
|
||||
usually at the expense of lower image quality. Note that the only available HunyuanVideo model is
|
||||
CFG-distilled, which means that traditional guidance between unconditional and conditional latent is
|
||||
not applied.
|
||||
num_videos_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
|
||||
generation deterministic.
|
||||
latents (`torch.Tensor`, *optional*):
|
||||
Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor is generated by sampling using the supplied random `generator`.
|
||||
prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs (prompt weighting). If not
|
||||
provided, text embeddings are generated from the `prompt` input argument.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`HunyuanVideoPipelineOutput`] instead of a plain tuple.
|
||||
attention_kwargs (`dict`, *optional*):
|
||||
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||
`self.processor` in
|
||||
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
clip_skip (`int`, *optional*):
|
||||
Number of layers to be skipped from CLIP while computing the prompt embeddings. A value of 1 means that
|
||||
the output of the pre-final layer will be used for computing the prompt embeddings.
|
||||
callback_on_step_end (`Callable`, `PipelineCallback`, `MultiPipelineCallbacks`, *optional*):
|
||||
A function or a subclass of `PipelineCallback` or `MultiPipelineCallbacks` that is called at the end of
|
||||
each denoising step during the inference. with the following arguments: `callback_on_step_end(self:
|
||||
DiffusionPipeline, step: int, timestep: int, callback_kwargs: Dict)`. `callback_kwargs` will include a
|
||||
list of all tensors as specified by `callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeline class.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~HunyuanVideoPipelineOutput`] or `tuple`:
|
||||
If `return_dict` is `True`, [`HunyuanVideoPipelineOutput`] is returned, otherwise a `tuple` is returned
|
||||
where the first element is a list with the generated images and the second element is a list of `bool`s
|
||||
indicating whether the corresponding generated image contains "not-safe-for-work" (nsfw) content.
|
||||
"""
|
||||
|
||||
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
|
||||
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(
|
||||
prompt,
|
||||
prompt_2,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds,
|
||||
callback_on_step_end_tensor_inputs,
|
||||
prompt_template,
|
||||
)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._attention_kwargs = attention_kwargs
|
||||
self._interrupt = False
|
||||
|
||||
device = self._execution_device
|
||||
|
||||
# 2. Define call parameters
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
# 3. Encode input prompt
|
||||
prompt_embeds, pooled_prompt_embeds, prompt_attention_mask = self.encode_prompt(
|
||||
prompt=prompt,
|
||||
prompt_2=prompt,
|
||||
prompt_template=prompt_template,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
prompt_embeds=prompt_embeds,
|
||||
pooled_prompt_embeds=pooled_prompt_embeds,
|
||||
prompt_attention_mask=prompt_attention_mask,
|
||||
device=device,
|
||||
max_sequence_length=max_sequence_length,
|
||||
)
|
||||
|
||||
transformer_dtype = self.transformer.dtype
|
||||
prompt_embeds = prompt_embeds.to(transformer_dtype)
|
||||
prompt_attention_mask = prompt_attention_mask.to(transformer_dtype)
|
||||
if pooled_prompt_embeds is not None:
|
||||
pooled_prompt_embeds = pooled_prompt_embeds.to(transformer_dtype)
|
||||
|
||||
# 4. Prepare timesteps
|
||||
sigmas = np.linspace(1.0, 0.0, num_inference_steps + 1)[:-1] if sigmas is None else sigmas
|
||||
timesteps, num_inference_steps = retrieve_timesteps(
|
||||
self.scheduler,
|
||||
num_inference_steps,
|
||||
device,
|
||||
sigmas=sigmas,
|
||||
)
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = self.transformer.config.in_channels
|
||||
num_latent_frames = (num_frames - 1) // self.vae_scale_factor_temporal + 1
|
||||
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
num_latent_frames,
|
||||
torch.float32,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
# check sequence_parallel
|
||||
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
|
||||
if get_sequence_parallel_state():
|
||||
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
|
||||
latents = latents[:, :, rank, :, :, :]
|
||||
|
||||
# 6. Prepare guidance condition
|
||||
guidance = torch.tensor([guidance_scale] * latents.shape[0], dtype=transformer_dtype, device=device) * 1000.0
|
||||
|
||||
# 7. Denoising loop
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
latent_model_input = latents.to(transformer_dtype)
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latents.shape[0]).to(latents.dtype)
|
||||
if pooled_prompt_embeds.shape[-1] != prompt_embeds.shape[-1]:
|
||||
pooled_prompt_embeds_padding = F.pad(
|
||||
pooled_prompt_embeds,
|
||||
(0, prompt_embeds.shape[2] - pooled_prompt_embeds.shape[1]),
|
||||
value=0,
|
||||
).unsqueeze(1)
|
||||
encoder_hidden_states = torch.cat([pooled_prompt_embeds_padding, prompt_embeds], dim=1)
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=encoder_hidden_states, # [1, 257, 4096]
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
guidance=guidance,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
for k in callback_on_step_end_tensor_inputs:
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
latents = all_gather(latents, dim=2)
|
||||
|
||||
if not output_type == "latent":
|
||||
latents = latents.to(self.vae.dtype) / self.vae.config.scaling_factor
|
||||
video = self.vae.decode(latents, return_dict=False)[0]
|
||||
video = self.video_processor.postprocess_video(video, output_type=output_type)
|
||||
else:
|
||||
video = latents
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return (video, )
|
||||
|
||||
return HunyuanVideoPipelineOutput(frames=video)
|
||||
@@ -1,14 +1,19 @@
|
||||
import argparse
|
||||
import os
|
||||
|
||||
import torch
|
||||
import argparse
|
||||
from safetensors.torch import save_file
|
||||
import os
|
||||
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--diffusers_path", required=True, type=str)
|
||||
parser.add_argument("--transformer_path", type=str, default=None, help="Path to save transformer model")
|
||||
parser.add_argument("--vae_encoder_path", type=str, default=None, help="Path to save VAE encoder model")
|
||||
parser.add_argument("--vae_decoder_path", type=str, default=None, help="Path to save VAE decoder model")
|
||||
parser.add_argument(
|
||||
"--transformer_path", type=str, default=None, help="Path to save transformer model"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae_encoder_path", type=str, default=None, help="Path to save VAE encoder model"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae_decoder_path", type=str, default=None, help="Path to save VAE decoder model"
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
@@ -30,22 +35,50 @@ def convert_diffusers_transformer_to_mochi(state_dict):
|
||||
new_state_dict = {}
|
||||
|
||||
# Convert patch_embed
|
||||
new_state_dict["x_embedder.proj.weight"] = original_state_dict.pop("patch_embed.proj.weight")
|
||||
new_state_dict["x_embedder.proj.bias"] = original_state_dict.pop("patch_embed.proj.bias")
|
||||
new_state_dict["x_embedder.proj.weight"] = original_state_dict.pop(
|
||||
"patch_embed.proj.weight"
|
||||
)
|
||||
new_state_dict["x_embedder.proj.bias"] = original_state_dict.pop(
|
||||
"patch_embed.proj.bias"
|
||||
)
|
||||
|
||||
# Convert time_embed
|
||||
new_state_dict["t_embedder.mlp.0.weight"] = original_state_dict.pop("time_embed.timestep_embedder.linear_1.weight")
|
||||
new_state_dict["t_embedder.mlp.0.bias"] = original_state_dict.pop("time_embed.timestep_embedder.linear_1.bias")
|
||||
new_state_dict["t_embedder.mlp.2.weight"] = original_state_dict.pop("time_embed.timestep_embedder.linear_2.weight")
|
||||
new_state_dict["t_embedder.mlp.2.bias"] = original_state_dict.pop("time_embed.timestep_embedder.linear_2.bias")
|
||||
new_state_dict["t5_y_embedder.to_kv.weight"] = original_state_dict.pop("time_embed.pooler.to_kv.weight")
|
||||
new_state_dict["t5_y_embedder.to_kv.bias"] = original_state_dict.pop("time_embed.pooler.to_kv.bias")
|
||||
new_state_dict["t5_y_embedder.to_q.weight"] = original_state_dict.pop("time_embed.pooler.to_q.weight")
|
||||
new_state_dict["t5_y_embedder.to_q.bias"] = original_state_dict.pop("time_embed.pooler.to_q.bias")
|
||||
new_state_dict["t5_y_embedder.to_out.weight"] = original_state_dict.pop("time_embed.pooler.to_out.weight")
|
||||
new_state_dict["t5_y_embedder.to_out.bias"] = original_state_dict.pop("time_embed.pooler.to_out.bias")
|
||||
new_state_dict["t5_yproj.weight"] = original_state_dict.pop("time_embed.caption_proj.weight")
|
||||
new_state_dict["t5_yproj.bias"] = original_state_dict.pop("time_embed.caption_proj.bias")
|
||||
new_state_dict["t_embedder.mlp.0.weight"] = original_state_dict.pop(
|
||||
"time_embed.timestep_embedder.linear_1.weight"
|
||||
)
|
||||
new_state_dict["t_embedder.mlp.0.bias"] = original_state_dict.pop(
|
||||
"time_embed.timestep_embedder.linear_1.bias"
|
||||
)
|
||||
new_state_dict["t_embedder.mlp.2.weight"] = original_state_dict.pop(
|
||||
"time_embed.timestep_embedder.linear_2.weight"
|
||||
)
|
||||
new_state_dict["t_embedder.mlp.2.bias"] = original_state_dict.pop(
|
||||
"time_embed.timestep_embedder.linear_2.bias"
|
||||
)
|
||||
new_state_dict["t5_y_embedder.to_kv.weight"] = original_state_dict.pop(
|
||||
"time_embed.pooler.to_kv.weight"
|
||||
)
|
||||
new_state_dict["t5_y_embedder.to_kv.bias"] = original_state_dict.pop(
|
||||
"time_embed.pooler.to_kv.bias"
|
||||
)
|
||||
new_state_dict["t5_y_embedder.to_q.weight"] = original_state_dict.pop(
|
||||
"time_embed.pooler.to_q.weight"
|
||||
)
|
||||
new_state_dict["t5_y_embedder.to_q.bias"] = original_state_dict.pop(
|
||||
"time_embed.pooler.to_q.bias"
|
||||
)
|
||||
new_state_dict["t5_y_embedder.to_out.weight"] = original_state_dict.pop(
|
||||
"time_embed.pooler.to_out.weight"
|
||||
)
|
||||
new_state_dict["t5_y_embedder.to_out.bias"] = original_state_dict.pop(
|
||||
"time_embed.pooler.to_out.bias"
|
||||
)
|
||||
new_state_dict["t5_yproj.weight"] = original_state_dict.pop(
|
||||
"time_embed.caption_proj.weight"
|
||||
)
|
||||
new_state_dict["t5_yproj.bias"] = original_state_dict.pop(
|
||||
"time_embed.caption_proj.bias"
|
||||
)
|
||||
|
||||
# Convert transformer blocks
|
||||
num_layers = 48
|
||||
@@ -54,19 +87,27 @@ def convert_diffusers_transformer_to_mochi(state_dict):
|
||||
new_prefix = f"blocks.{i}."
|
||||
|
||||
# norm1
|
||||
new_state_dict[new_prefix + "mod_x.weight"] = original_state_dict.pop(block_prefix + "norm1.linear.weight")
|
||||
new_state_dict[new_prefix + "mod_x.bias"] = original_state_dict.pop(block_prefix + "norm1.linear.bias")
|
||||
new_state_dict[new_prefix + "mod_x.weight"] = original_state_dict.pop(
|
||||
block_prefix + "norm1.linear.weight"
|
||||
)
|
||||
new_state_dict[new_prefix + "mod_x.bias"] = original_state_dict.pop(
|
||||
block_prefix + "norm1.linear.bias"
|
||||
)
|
||||
|
||||
if i < num_layers - 1:
|
||||
new_state_dict[new_prefix + "mod_y.weight"] = original_state_dict.pop(block_prefix +
|
||||
"norm1_context.linear.weight")
|
||||
new_state_dict[new_prefix + "mod_y.bias"] = original_state_dict.pop(block_prefix +
|
||||
"norm1_context.linear.bias")
|
||||
new_state_dict[new_prefix + "mod_y.weight"] = original_state_dict.pop(
|
||||
block_prefix + "norm1_context.linear.weight"
|
||||
)
|
||||
new_state_dict[new_prefix + "mod_y.bias"] = original_state_dict.pop(
|
||||
block_prefix + "norm1_context.linear.bias"
|
||||
)
|
||||
else:
|
||||
new_state_dict[new_prefix + "mod_y.weight"] = original_state_dict.pop(block_prefix +
|
||||
"norm1_context.linear_1.weight")
|
||||
new_state_dict[new_prefix + "mod_y.bias"] = original_state_dict.pop(block_prefix +
|
||||
"norm1_context.linear_1.bias")
|
||||
new_state_dict[new_prefix + "mod_y.weight"] = original_state_dict.pop(
|
||||
block_prefix + "norm1_context.linear_1.weight"
|
||||
)
|
||||
new_state_dict[new_prefix + "mod_y.bias"] = original_state_dict.pop(
|
||||
block_prefix + "norm1_context.linear_1.bias"
|
||||
)
|
||||
|
||||
# Visual attention
|
||||
q = original_state_dict.pop(block_prefix + "attn1.to_q.weight")
|
||||
@@ -75,13 +116,18 @@ def convert_diffusers_transformer_to_mochi(state_dict):
|
||||
qkv_weight = torch.cat([q, k, v], dim=0)
|
||||
new_state_dict[new_prefix + "attn.qkv_x.weight"] = qkv_weight
|
||||
|
||||
new_state_dict[new_prefix + "attn.q_norm_x.weight"] = original_state_dict.pop(block_prefix +
|
||||
"attn1.norm_q.weight")
|
||||
new_state_dict[new_prefix + "attn.k_norm_x.weight"] = original_state_dict.pop(block_prefix +
|
||||
"attn1.norm_k.weight")
|
||||
new_state_dict[new_prefix + "attn.proj_x.weight"] = original_state_dict.pop(block_prefix +
|
||||
"attn1.to_out.0.weight")
|
||||
new_state_dict[new_prefix + "attn.proj_x.bias"] = original_state_dict.pop(block_prefix + "attn1.to_out.0.bias")
|
||||
new_state_dict[new_prefix + "attn.q_norm_x.weight"] = original_state_dict.pop(
|
||||
block_prefix + "attn1.norm_q.weight"
|
||||
)
|
||||
new_state_dict[new_prefix + "attn.k_norm_x.weight"] = original_state_dict.pop(
|
||||
block_prefix + "attn1.norm_k.weight"
|
||||
)
|
||||
new_state_dict[new_prefix + "attn.proj_x.weight"] = original_state_dict.pop(
|
||||
block_prefix + "attn1.to_out.0.weight"
|
||||
)
|
||||
new_state_dict[new_prefix + "attn.proj_x.bias"] = original_state_dict.pop(
|
||||
block_prefix + "attn1.to_out.0.bias"
|
||||
)
|
||||
|
||||
# Context attention
|
||||
q = original_state_dict.pop(block_prefix + "attn1.add_q_proj.weight")
|
||||
@@ -90,31 +136,45 @@ def convert_diffusers_transformer_to_mochi(state_dict):
|
||||
qkv_weight = torch.cat([q, k, v], dim=0)
|
||||
new_state_dict[new_prefix + "attn.qkv_y.weight"] = qkv_weight
|
||||
|
||||
new_state_dict[new_prefix + "attn.q_norm_y.weight"] = original_state_dict.pop(block_prefix +
|
||||
"attn1.norm_added_q.weight")
|
||||
new_state_dict[new_prefix + "attn.k_norm_y.weight"] = original_state_dict.pop(block_prefix +
|
||||
"attn1.norm_added_k.weight")
|
||||
new_state_dict[new_prefix + "attn.q_norm_y.weight"] = original_state_dict.pop(
|
||||
block_prefix + "attn1.norm_added_q.weight"
|
||||
)
|
||||
new_state_dict[new_prefix + "attn.k_norm_y.weight"] = original_state_dict.pop(
|
||||
block_prefix + "attn1.norm_added_k.weight"
|
||||
)
|
||||
if i < num_layers - 1:
|
||||
new_state_dict[new_prefix + "attn.proj_y.weight"] = original_state_dict.pop(block_prefix +
|
||||
"attn1.to_add_out.weight")
|
||||
new_state_dict[new_prefix + "attn.proj_y.bias"] = original_state_dict.pop(block_prefix +
|
||||
"attn1.to_add_out.bias")
|
||||
new_state_dict[new_prefix + "attn.proj_y.weight"] = original_state_dict.pop(
|
||||
block_prefix + "attn1.to_add_out.weight"
|
||||
)
|
||||
new_state_dict[new_prefix + "attn.proj_y.bias"] = original_state_dict.pop(
|
||||
block_prefix + "attn1.to_add_out.bias"
|
||||
)
|
||||
|
||||
# MLP
|
||||
new_state_dict[new_prefix + "mlp_x.w1.weight"] = reverse_proj_gate(
|
||||
original_state_dict.pop(block_prefix + "ff.net.0.proj.weight"))
|
||||
new_state_dict[new_prefix + "mlp_x.w2.weight"] = original_state_dict.pop(block_prefix + "ff.net.2.weight")
|
||||
original_state_dict.pop(block_prefix + "ff.net.0.proj.weight")
|
||||
)
|
||||
new_state_dict[new_prefix + "mlp_x.w2.weight"] = original_state_dict.pop(
|
||||
block_prefix + "ff.net.2.weight"
|
||||
)
|
||||
if i < num_layers - 1:
|
||||
new_state_dict[new_prefix + "mlp_y.w1.weight"] = reverse_proj_gate(
|
||||
original_state_dict.pop(block_prefix + "ff_context.net.0.proj.weight"))
|
||||
new_state_dict[new_prefix + "mlp_y.w2.weight"] = original_state_dict.pop(block_prefix +
|
||||
"ff_context.net.2.weight")
|
||||
original_state_dict.pop(block_prefix + "ff_context.net.0.proj.weight")
|
||||
)
|
||||
new_state_dict[new_prefix + "mlp_y.w2.weight"] = original_state_dict.pop(
|
||||
block_prefix + "ff_context.net.2.weight"
|
||||
)
|
||||
|
||||
# Output layers
|
||||
new_state_dict["final_layer.mod.weight"] = reverse_scale_shift(original_state_dict.pop("norm_out.linear.weight"),
|
||||
dim=0)
|
||||
new_state_dict["final_layer.mod.bias"] = reverse_scale_shift(original_state_dict.pop("norm_out.linear.bias"), dim=0)
|
||||
new_state_dict["final_layer.linear.weight"] = original_state_dict.pop("proj_out.weight")
|
||||
new_state_dict["final_layer.mod.weight"] = reverse_scale_shift(
|
||||
original_state_dict.pop("norm_out.linear.weight"), dim=0
|
||||
)
|
||||
new_state_dict["final_layer.mod.bias"] = reverse_scale_shift(
|
||||
original_state_dict.pop("norm_out.linear.bias"), dim=0
|
||||
)
|
||||
new_state_dict["final_layer.linear.weight"] = original_state_dict.pop(
|
||||
"proj_out.weight"
|
||||
)
|
||||
new_state_dict["final_layer.linear.bias"] = original_state_dict.pop("proj_out.bias")
|
||||
|
||||
new_state_dict["pos_frequencies"] = original_state_dict.pop("pos_frequencies")
|
||||
@@ -132,89 +192,155 @@ def convert_diffusers_vae_to_mochi(state_dict):
|
||||
# Convert encoder
|
||||
prefix = "encoder."
|
||||
|
||||
encoder_state_dict["layers.0.weight"] = original_state_dict.pop(f"{prefix}proj_in.weight")
|
||||
encoder_state_dict["layers.0.bias"] = original_state_dict.pop(f"{prefix}proj_in.bias")
|
||||
encoder_state_dict["layers.0.weight"] = original_state_dict.pop(
|
||||
f"{prefix}proj_in.weight"
|
||||
)
|
||||
encoder_state_dict["layers.0.bias"] = original_state_dict.pop(
|
||||
f"{prefix}proj_in.bias"
|
||||
)
|
||||
|
||||
# Convert block_in
|
||||
for i in range(3):
|
||||
encoder_state_dict[f"layers.{i+1}.stack.0.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.weight")
|
||||
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.weight"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+1}.stack.0.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.bias")
|
||||
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.bias"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+1}.stack.2.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.conv1.conv.weight")
|
||||
f"{prefix}block_in.resnets.{i}.conv1.conv.weight"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+1}.stack.2.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.conv1.conv.bias")
|
||||
f"{prefix}block_in.resnets.{i}.conv1.conv.bias"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+1}.stack.3.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.weight")
|
||||
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.weight"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+1}.stack.3.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.bias")
|
||||
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.bias"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+1}.stack.5.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.conv2.conv.weight")
|
||||
f"{prefix}block_in.resnets.{i}.conv2.conv.weight"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+1}.stack.5.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.conv2.conv.bias")
|
||||
f"{prefix}block_in.resnets.{i}.conv2.conv.bias"
|
||||
)
|
||||
|
||||
# Convert down_blocks
|
||||
down_block_layers = [3, 4, 6]
|
||||
for block in range(3):
|
||||
encoder_state_dict[f"layers.{block+4}.layers.0.weight"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.conv_in.conv.weight")
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.0.weight"
|
||||
] = original_state_dict.pop(f"{prefix}down_blocks.{block}.conv_in.conv.weight")
|
||||
encoder_state_dict[f"layers.{block+4}.layers.0.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.conv_in.conv.bias")
|
||||
f"{prefix}down_blocks.{block}.conv_in.conv.bias"
|
||||
)
|
||||
|
||||
for i in range(down_block_layers[block]):
|
||||
# Convert resnets
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.0.weight"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.weight")
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.0.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.bias")
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.2.weight"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.weight")
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.2.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.bias")
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.3.weight"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.weight")
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.3.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.bias")
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.5.weight"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.weight")
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.stack.5.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.bias")
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.0.weight"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.weight"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.0.bias"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.norm1.norm_layer.bias"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.2.weight"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.weight"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.2.bias"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.conv1.conv.bias"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.3.weight"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.weight"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.3.bias"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.norm2.norm_layer.bias"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.5.weight"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.weight"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.stack.5.bias"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.resnets.{i}.conv2.conv.bias"
|
||||
)
|
||||
|
||||
# Convert attentions
|
||||
q = original_state_dict.pop(f"{prefix}down_blocks.{block}.attentions.{i}.to_q.weight")
|
||||
k = original_state_dict.pop(f"{prefix}down_blocks.{block}.attentions.{i}.to_k.weight")
|
||||
v = original_state_dict.pop(f"{prefix}down_blocks.{block}.attentions.{i}.to_v.weight")
|
||||
q = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.attentions.{i}.to_q.weight"
|
||||
)
|
||||
k = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.attentions.{i}.to_k.weight"
|
||||
)
|
||||
v = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.attentions.{i}.to_v.weight"
|
||||
)
|
||||
qkv_weight = torch.cat([q, k, v], dim=0)
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.attn.qkv.weight"] = qkv_weight
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.attn_block.attn.qkv.weight"
|
||||
] = qkv_weight
|
||||
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.weight"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.attentions.{i}.to_out.0.weight")
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.attentions.{i}.to_out.0.bias")
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.norm.weight"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.norms.{i}.norm_layer.weight")
|
||||
encoder_state_dict[f"layers.{block+4}.layers.{i+1}.attn_block.norm.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.norms.{i}.norm_layer.bias")
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.weight"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.attentions.{i}.to_out.0.weight"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.attn_block.attn.out.bias"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.attentions.{i}.to_out.0.bias"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.attn_block.norm.weight"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.norms.{i}.norm_layer.weight"
|
||||
)
|
||||
encoder_state_dict[
|
||||
f"layers.{block+4}.layers.{i+1}.attn_block.norm.bias"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.norms.{i}.norm_layer.bias"
|
||||
)
|
||||
|
||||
# Convert block_out
|
||||
for i in range(3):
|
||||
encoder_state_dict[f"layers.{i+7}.stack.0.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.weight")
|
||||
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.weight"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+7}.stack.0.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.bias")
|
||||
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.bias"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+7}.stack.2.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.conv1.conv.weight")
|
||||
f"{prefix}block_out.resnets.{i}.conv1.conv.weight"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+7}.stack.2.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.conv1.conv.bias")
|
||||
f"{prefix}block_out.resnets.{i}.conv1.conv.bias"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+7}.stack.3.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.weight")
|
||||
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.weight"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+7}.stack.3.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.bias")
|
||||
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.bias"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+7}.stack.5.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.conv2.conv.weight")
|
||||
f"{prefix}block_out.resnets.{i}.conv2.conv.weight"
|
||||
)
|
||||
encoder_state_dict[f"layers.{i+7}.stack.5.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.conv2.conv.bias")
|
||||
f"{prefix}block_out.resnets.{i}.conv2.conv.bias"
|
||||
)
|
||||
|
||||
q = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_q.weight")
|
||||
k = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_k.weight")
|
||||
@@ -222,92 +348,152 @@ def convert_diffusers_vae_to_mochi(state_dict):
|
||||
qkv_weight = torch.cat([q, k, v], dim=0)
|
||||
encoder_state_dict[f"layers.{i+7}.attn_block.attn.qkv.weight"] = qkv_weight
|
||||
|
||||
encoder_state_dict[f"layers.{i+7}.attn_block.attn.out.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.attentions.{i}.to_out.0.weight")
|
||||
encoder_state_dict[f"layers.{i+7}.attn_block.attn.out.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.attentions.{i}.to_out.0.bias")
|
||||
encoder_state_dict[f"layers.{i+7}.attn_block.norm.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.norms.{i}.norm_layer.weight")
|
||||
encoder_state_dict[f"layers.{i+7}.attn_block.norm.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.norms.{i}.norm_layer.bias")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+7}.attn_block.attn.out.weight"
|
||||
] = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_out.0.weight")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+7}.attn_block.attn.out.bias"
|
||||
] = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_out.0.bias")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+7}.attn_block.norm.weight"
|
||||
] = original_state_dict.pop(f"{prefix}block_out.norms.{i}.norm_layer.weight")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+7}.attn_block.norm.bias"
|
||||
] = original_state_dict.pop(f"{prefix}block_out.norms.{i}.norm_layer.bias")
|
||||
|
||||
# Convert output layers
|
||||
encoder_state_dict["output_norm.weight"] = original_state_dict.pop(f"{prefix}norm_out.norm_layer.weight")
|
||||
encoder_state_dict["output_norm.bias"] = original_state_dict.pop(f"{prefix}norm_out.norm_layer.bias")
|
||||
encoder_state_dict["output_proj.weight"] = original_state_dict.pop(f"{prefix}proj_out.weight")
|
||||
encoder_state_dict["output_norm.weight"] = original_state_dict.pop(
|
||||
f"{prefix}norm_out.norm_layer.weight"
|
||||
)
|
||||
encoder_state_dict["output_norm.bias"] = original_state_dict.pop(
|
||||
f"{prefix}norm_out.norm_layer.bias"
|
||||
)
|
||||
encoder_state_dict["output_proj.weight"] = original_state_dict.pop(
|
||||
f"{prefix}proj_out.weight"
|
||||
)
|
||||
|
||||
# Convert decoder
|
||||
prefix = "decoder."
|
||||
|
||||
decoder_state_dict["blocks.0.0.weight"] = original_state_dict.pop(f"{prefix}conv_in.weight")
|
||||
decoder_state_dict["blocks.0.0.bias"] = original_state_dict.pop(f"{prefix}conv_in.bias")
|
||||
decoder_state_dict["blocks.0.0.weight"] = original_state_dict.pop(
|
||||
f"{prefix}conv_in.weight"
|
||||
)
|
||||
decoder_state_dict["blocks.0.0.bias"] = original_state_dict.pop(
|
||||
f"{prefix}conv_in.bias"
|
||||
)
|
||||
|
||||
# Convert block_in
|
||||
for i in range(3):
|
||||
decoder_state_dict[f"blocks.0.{i+1}.stack.0.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.weight")
|
||||
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.weight"
|
||||
)
|
||||
decoder_state_dict[f"blocks.0.{i+1}.stack.0.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.bias")
|
||||
f"{prefix}block_in.resnets.{i}.norm1.norm_layer.bias"
|
||||
)
|
||||
decoder_state_dict[f"blocks.0.{i+1}.stack.2.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.conv1.conv.weight")
|
||||
f"{prefix}block_in.resnets.{i}.conv1.conv.weight"
|
||||
)
|
||||
decoder_state_dict[f"blocks.0.{i+1}.stack.2.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.conv1.conv.bias")
|
||||
f"{prefix}block_in.resnets.{i}.conv1.conv.bias"
|
||||
)
|
||||
decoder_state_dict[f"blocks.0.{i+1}.stack.3.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.weight")
|
||||
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.weight"
|
||||
)
|
||||
decoder_state_dict[f"blocks.0.{i+1}.stack.3.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.bias")
|
||||
f"{prefix}block_in.resnets.{i}.norm2.norm_layer.bias"
|
||||
)
|
||||
decoder_state_dict[f"blocks.0.{i+1}.stack.5.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.conv2.conv.weight")
|
||||
f"{prefix}block_in.resnets.{i}.conv2.conv.weight"
|
||||
)
|
||||
decoder_state_dict[f"blocks.0.{i+1}.stack.5.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.conv2.conv.bias")
|
||||
f"{prefix}block_in.resnets.{i}.conv2.conv.bias"
|
||||
)
|
||||
|
||||
# Convert up_blocks
|
||||
up_block_layers = [6, 4, 3]
|
||||
for block in range(3):
|
||||
for i in range(up_block_layers[block]):
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.0.weight"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.weight")
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.0.bias"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.bias")
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.2.weight"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.weight")
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.2.bias"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.bias")
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.3.weight"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.weight")
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.3.bias"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.bias")
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.5.weight"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.weight")
|
||||
decoder_state_dict[f"blocks.{block+1}.blocks.{i}.stack.5.bias"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.bias")
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.0.weight"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.weight"
|
||||
)
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.0.bias"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.norm1.norm_layer.bias"
|
||||
)
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.2.weight"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.weight"
|
||||
)
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.2.bias"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.conv1.conv.bias"
|
||||
)
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.3.weight"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.weight"
|
||||
)
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.3.bias"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.norm2.norm_layer.bias"
|
||||
)
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.5.weight"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.weight"
|
||||
)
|
||||
decoder_state_dict[
|
||||
f"blocks.{block+1}.blocks.{i}.stack.5.bias"
|
||||
] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.resnets.{i}.conv2.conv.bias"
|
||||
)
|
||||
decoder_state_dict[f"blocks.{block+1}.proj.weight"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.proj.weight")
|
||||
f"{prefix}up_blocks.{block}.proj.weight"
|
||||
)
|
||||
decoder_state_dict[f"blocks.{block+1}.proj.bias"] = original_state_dict.pop(
|
||||
f"{prefix}up_blocks.{block}.proj.bias")
|
||||
f"{prefix}up_blocks.{block}.proj.bias"
|
||||
)
|
||||
|
||||
# Convert block_out
|
||||
for i in range(3):
|
||||
decoder_state_dict[f"blocks.4.{i}.stack.0.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.weight")
|
||||
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.weight"
|
||||
)
|
||||
decoder_state_dict[f"blocks.4.{i}.stack.0.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.bias")
|
||||
f"{prefix}block_out.resnets.{i}.norm1.norm_layer.bias"
|
||||
)
|
||||
decoder_state_dict[f"blocks.4.{i}.stack.2.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.conv1.conv.weight")
|
||||
f"{prefix}block_out.resnets.{i}.conv1.conv.weight"
|
||||
)
|
||||
decoder_state_dict[f"blocks.4.{i}.stack.2.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.conv1.conv.bias")
|
||||
f"{prefix}block_out.resnets.{i}.conv1.conv.bias"
|
||||
)
|
||||
decoder_state_dict[f"blocks.4.{i}.stack.3.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.weight")
|
||||
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.weight"
|
||||
)
|
||||
decoder_state_dict[f"blocks.4.{i}.stack.3.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.bias")
|
||||
f"{prefix}block_out.resnets.{i}.norm2.norm_layer.bias"
|
||||
)
|
||||
decoder_state_dict[f"blocks.4.{i}.stack.5.weight"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.conv2.conv.weight")
|
||||
f"{prefix}block_out.resnets.{i}.conv2.conv.weight"
|
||||
)
|
||||
decoder_state_dict[f"blocks.4.{i}.stack.5.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.conv2.conv.bias")
|
||||
f"{prefix}block_out.resnets.{i}.conv2.conv.bias"
|
||||
)
|
||||
|
||||
# Convert output layers
|
||||
decoder_state_dict["output_proj.weight"] = original_state_dict.pop(f"{prefix}proj_out.weight")
|
||||
decoder_state_dict["output_proj.bias"] = original_state_dict.pop(f"{prefix}proj_out.bias")
|
||||
decoder_state_dict["output_proj.weight"] = original_state_dict.pop(
|
||||
f"{prefix}proj_out.weight"
|
||||
)
|
||||
decoder_state_dict["output_proj.bias"] = original_state_dict.pop(
|
||||
f"{prefix}proj_out.bias"
|
||||
)
|
||||
|
||||
return encoder_state_dict, decoder_state_dict
|
||||
|
||||
@@ -333,8 +519,10 @@ def main(args):
|
||||
transformer_path = ensure_safetensors_extension(args.transformer_path)
|
||||
ensure_directory_exists(transformer_path)
|
||||
|
||||
print("Converting transformer model...")
|
||||
transformer_state_dict = convert_diffusers_transformer_to_mochi(pipe.transformer.state_dict())
|
||||
print(f"Converting transformer model...")
|
||||
transformer_state_dict = convert_diffusers_transformer_to_mochi(
|
||||
pipe.transformer.state_dict()
|
||||
)
|
||||
save_file(transformer_state_dict, transformer_path)
|
||||
print(f"Saved transformer to {transformer_path}")
|
||||
|
||||
@@ -345,8 +533,10 @@ def main(args):
|
||||
ensure_directory_exists(encoder_path)
|
||||
ensure_directory_exists(decoder_path)
|
||||
|
||||
print("Converting VAE models...")
|
||||
encoder_state_dict, decoder_state_dict = convert_diffusers_vae_to_mochi(pipe.vae.state_dict())
|
||||
print(f"Converting VAE models...")
|
||||
encoder_state_dict, decoder_state_dict = convert_diffusers_vae_to_mochi(
|
||||
pipe.vae.state_dict()
|
||||
)
|
||||
|
||||
save_file(encoder_state_dict, encoder_path)
|
||||
print(f"Saved VAE encoder to {encoder_path}")
|
||||
@@ -354,7 +544,9 @@ def main(args):
|
||||
save_file(decoder_state_dict, decoder_path)
|
||||
print(f"Saved VAE decoder to {decoder_path}")
|
||||
elif args.vae_encoder_path or args.vae_decoder_path:
|
||||
print("Warning: Both VAE encoder and decoder paths must be specified to convert VAE models.")
|
||||
print(
|
||||
"Warning: Both VAE encoder and decoder paths must be specified to convert VAE models."
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user