Compare commits
31
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e8986b0b79 | ||
|
|
15553f7706 | ||
|
|
60eeea50bb | ||
|
|
927b3a40b9 | ||
|
|
c64f826ae2 | ||
|
|
55c1040f0b | ||
|
|
8a3e7aa761 | ||
|
|
4324c1c21d | ||
|
|
2c342ee37f | ||
|
|
708201f531 | ||
|
|
1fee098f10 | ||
|
|
8a77cf22c9 | ||
|
|
d869d90d12 | ||
|
|
554ee17de5 | ||
|
|
0be4fc62c9 | ||
|
|
1e08893546 | ||
|
|
09ab452610 | ||
|
|
e768b5ec5b | ||
|
|
59ec42f40e | ||
|
|
5ae5b247b3 | ||
|
|
ead6c62be4 | ||
|
|
6805eaa06c | ||
|
|
c39a15551c | ||
|
|
e6dda263b0 | ||
|
|
f9482d113c | ||
|
|
a3ec969397 | ||
|
|
76a12cc8a1 | ||
|
|
ac490399c6 | ||
|
|
9ea39cee57 | ||
|
|
52e6e612a2 | ||
|
|
9aebc4ada1 |
@@ -0,0 +1 @@
|
||||
blank_issues_enabled: false
|
||||
@@ -0,0 +1,239 @@
|
||||
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
|
||||
}
|
||||
|
||||
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=1)
|
||||
|
||||
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()
|
||||
@@ -0,0 +1,90 @@
|
||||
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,45 +0,0 @@
|
||||
name: codespell
|
||||
|
||||
on:
|
||||
# Trigger the workflow on push or pull request,
|
||||
# but only for the main branch
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- "**/*.md"
|
||||
- "**/*.rst"
|
||||
- pyproject.toml
|
||||
- requirements-lint.txt
|
||||
- .github/workflows/codespell.yml
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- "**/*.md"
|
||||
- "**/*.rst"
|
||||
- pyproject.toml
|
||||
- requirements-lint.txt
|
||||
- .github/workflows/codespell.yml
|
||||
|
||||
jobs:
|
||||
codespell:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v3
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: '3.12' # or any version you need
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements-lint.txt
|
||||
- name: Spelling check with codespell
|
||||
run: |
|
||||
# Refer to the above environment variable here
|
||||
codespell --toml pyproject.toml $CODESPELL_EXCLUDES
|
||||
@@ -0,0 +1,70 @@
|
||||
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/
|
||||
@@ -0,0 +1,17 @@
|
||||
{
|
||||
"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
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,16 @@
|
||||
{
|
||||
"problemMatcher": [
|
||||
{
|
||||
"owner": "mypy",
|
||||
"pattern": [
|
||||
{
|
||||
"regexp": "^(.+):(\\d+):\\s(error|warning):\\s(.+)$",
|
||||
"file": 1,
|
||||
"line": 2,
|
||||
"severity": 3,
|
||||
"message": 4
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
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
|
||||
@@ -0,0 +1,18 @@
|
||||
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,50 +0,0 @@
|
||||
name: ruff
|
||||
|
||||
on:
|
||||
# Trigger the workflow on push or pull request,
|
||||
# but only for the main branch
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- pyproject.toml
|
||||
- requirements-lint.txt
|
||||
- .github/workflows/matchers/ruff.json
|
||||
- .github/workflows/ruff.yml
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
# This workflow is only relevant when one of the following files changes.
|
||||
# However, we have github configured to expect and require this workflow
|
||||
# to run and pass before github with auto-merge a pull request. Until github
|
||||
# allows more flexible auto-merge policy, we can just run this on every PR.
|
||||
# It doesn't take that long to run, anyway.
|
||||
#paths:
|
||||
# - "**/*.py"
|
||||
# - pyproject.toml
|
||||
# - requirements-lint.txt
|
||||
# - .github/workflows/matchers/ruff.json
|
||||
# - .github/workflows/ruff.yml
|
||||
|
||||
jobs:
|
||||
ruff:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v3
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: '3.12' # or any version you need
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements-lint.txt
|
||||
- name: Analysing the code with ruff
|
||||
run: |
|
||||
ruff check .
|
||||
- name: Run isort
|
||||
run: |
|
||||
isort . --check-only
|
||||
@@ -0,0 +1,221 @@
|
||||
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/
|
||||
@@ -11,10 +11,10 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v3
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v4
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.12' # or any version you need
|
||||
|
||||
@@ -25,6 +25,7 @@ jobs:
|
||||
pip install packaging ninja
|
||||
pip install -e .
|
||||
pip install pytest
|
||||
|
||||
- name: Run Pytest
|
||||
run: |
|
||||
pytest
|
||||
pytest --ignore csrc/sliding_tile_attention/test
|
||||
|
||||
@@ -1,38 +0,0 @@
|
||||
name: yapf
|
||||
|
||||
on:
|
||||
# Trigger the workflow on push or pull request,
|
||||
# but only for the main branch
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- .github/workflows/yapf.yml
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- .github/workflows/yapf.yml
|
||||
|
||||
jobs:
|
||||
yapf:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v3
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: '3.12' # or any version you need
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install yapf==0.32.0
|
||||
pip install toml==0.10.2
|
||||
- name: Running yapf
|
||||
run: |
|
||||
yapf --diff --recursive .
|
||||
+9
-26
@@ -1,6 +1,4 @@
|
||||
ucf101_stride4x4x4
|
||||
__pycache__
|
||||
*.mp4
|
||||
.ipynb_checkpoints
|
||||
*.pth
|
||||
UCF-101/
|
||||
@@ -8,40 +6,17 @@ 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/
|
||||
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/
|
||||
@@ -51,3 +26,11 @@ outputs_video
|
||||
sbatch.sh
|
||||
*.out
|
||||
env
|
||||
dist/
|
||||
*.o
|
||||
**/build/
|
||||
**.egg-info
|
||||
**.pyc
|
||||
**.egg
|
||||
**.txt
|
||||
**.json
|
||||
@@ -0,0 +1,3 @@
|
||||
[submodule "csrc/sliding_tile_attention/tk"]
|
||||
path = csrc/sliding_tile_attention/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
@@ -0,0 +1,80 @@
|
||||
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
|
||||
@@ -184,18 +184,4 @@
|
||||
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.
|
||||
|
||||
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.
|
||||
identification within third-party archives.
|
||||
@@ -4,16 +4,15 @@
|
||||
|
||||
FastVideo is a lightweight framework for accelerating large video diffusion models.
|
||||
|
||||
https://github.com/user-attachments/assets/064ac1d2-11ed-4a0c-955b-4d412a96ef30
|
||||
|
||||
|
||||
<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://discord.gg/REBzDQTWWt" target="_blank"> Discord </a> | 🕹️ <a href="https://replicate.com/lucataco/fast-hunyuan-video" target="_blank"> Replicate </a>
|
||||
</p>
|
||||
🤗 <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
|
||||
|
||||
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.
|
||||
@@ -22,62 +21,89 @@ FastVideo currently offers: (with more to come)
|
||||
|
||||
Dev in progress and highly experimental.
|
||||
|
||||
## 🎥 More Demos
|
||||
|
||||
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
|
||||
|
||||
## 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.
|
||||
```
|
||||
./env_setup.sh fastvideo
|
||||
# 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
|
||||
```
|
||||
|
||||
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
|
||||
```
|
||||
For more information about the VRAM requirements for BitsAndBytes quantization, please refer to the table below (timing measured on an H100 GPU):
|
||||
|
||||
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
|
||||
@@ -89,79 +115,99 @@ 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
|
||||
```
|
||||
|
||||
To launch the distillation process, use the following commands:
|
||||
|
||||
```
|
||||
bash scripts/distill/distill_hunyuan.sh # for hunyuan
|
||||
bash scripts/distill/distill_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
We also provide an optional script for distillation with adversarial loss, located at `fastvideo/distill_adv.py`. Although we tried adversarial loss, we did not observe significant improvements.
|
||||
## Finetune
|
||||
### ⚡ Full Finetune
|
||||
Ensure your data is prepared and preprocessed in the format specified in [data_preprocess.md](docs/data_preprocess.md). For convenience, we also provide a mochi preprocessed Black Myth Wukong data that can be downloaded directly:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Mochi-Black-Myth --local_dir=data/Mochi-Black-Myth --repo_type=dataset
|
||||
```
|
||||
|
||||
Download the original model weights as specified in [Distill Section](#-distill):
|
||||
|
||||
Then you can run the finetune with:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
### ⚡ Lora Finetune
|
||||
### ⚡ Lora Finetune
|
||||
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight --local_dir=data/Hunyuan-Black-Myth-Wukong-lora-weight --repo_type=model
|
||||
```
|
||||
|
||||
#### Minimum Hardware Requirement
|
||||
- 40 GB GPU memory each for 2 GPUs with lora.
|
||||
- 30 GB GPU memory each for 2 GPUs with CPU offload and lora.
|
||||
|
||||
- 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!
|
||||
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
|
||||
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
|
||||
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.
|
||||
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.
|
||||
|
||||
## 📑 Development Plan
|
||||
@@ -176,7 +222,7 @@ For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` opt
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please run `bash format.sh` before submitting a pull request.
|
||||
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.
|
||||
@@ -185,3 +231,27 @@ Run `pytest` to verify the data preprocessing, checkpoint saving, and sequence p
|
||||
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.
|
After Width: | Height: | Size: 751 KiB |
@@ -0,0 +1,2 @@
|
||||
recursive-include tk *
|
||||
include config.py
|
||||
@@ -0,0 +1,68 @@
|
||||
|
||||
|
||||
# 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.
|
||||
@@ -0,0 +1,15 @@
|
||||
### 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'
|
||||
@@ -0,0 +1,76 @@
|
||||
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"])
|
||||
@@ -0,0 +1,24 @@
|
||||
#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
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
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]
|
||||
@@ -0,0 +1,687 @@
|
||||
// # 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();
|
||||
}
|
||||
@@ -0,0 +1,151 @@
|
||||
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)
|
||||
@@ -0,0 +1,71 @@
|
||||
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
|
||||
@@ -0,0 +1,96 @@
|
||||
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
+1
Submodule csrc/sliding_tile_attention/tk added at 1719fb7264
+7
-22
@@ -23,9 +23,7 @@ def init_args():
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--seed", type=int, default=12345)
|
||||
parser.add_argument("--transformer_path", type=str, default=None)
|
||||
parser.add_argument("--scheduler_type",
|
||||
type=str,
|
||||
default="pcm_linear_quadratic")
|
||||
parser.add_argument("--scheduler_type", type=str, default="pcm_linear_quadratic")
|
||||
parser.add_argument("--lora_checkpoint_dir", type=str, default=None)
|
||||
parser.add_argument("--shift", type=float, default=8.0)
|
||||
parser.add_argument("--num_euler_timesteps", type=int, default=50)
|
||||
@@ -50,15 +48,11 @@ def load_model(args):
|
||||
)
|
||||
|
||||
if args.transformer_path:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(
|
||||
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:
|
||||
@@ -137,11 +131,7 @@ 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(
|
||||
@@ -164,8 +154,7 @@ 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,
|
||||
@@ -173,11 +162,7 @@ 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")
|
||||
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
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
|
||||
@@ -0,0 +1,24 @@
|
||||
# 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"
|
||||
@@ -0,0 +1,20 @@
|
||||
# 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:
|
||||
|
||||
```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,13 +33,16 @@ 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",
|
||||
@@ -62,7 +65,9 @@ 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.
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
@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
|
||||
@@ -0,0 +1,25 @@
|
||||
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
|
||||
@@ -0,0 +1,51 @@
|
||||
# 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.
|
||||
@@ -0,0 +1,8 @@
|
||||
.vertical-table-header th.head:not(.stub) {
|
||||
writing-mode: sideways-lr;
|
||||
white-space: nowrap;
|
||||
max-width: 0;
|
||||
p {
|
||||
margin: 0;
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
// 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));
|
||||
});
|
||||
});
|
||||
@@ -0,0 +1,39 @@
|
||||
<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>
|
||||
@@ -0,0 +1,259 @@
|
||||
# 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
|
||||
|
||||
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):
|
||||
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
|
||||
filename = info['module'].replace('.', '/')
|
||||
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__ # Get the class of the instance
|
||||
|
||||
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
|
||||
@@ -0,0 +1,242 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import itertools
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
|
||||
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)
|
||||
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: 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)]
|
||||
|
||||
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"
|
||||
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"
|
||||
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
|
||||
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())
|
||||
@@ -0,0 +1,10 @@
|
||||
# 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
|
||||
|
||||
:::
|
||||
@@ -0,0 +1,10 @@
|
||||
(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.
|
||||
@@ -0,0 +1,81 @@
|
||||
# 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
|
||||
:::
|
||||
|
||||
## Indices and tables
|
||||
|
||||
- {ref}`genindex`
|
||||
- {ref}`modindex`
|
||||
@@ -0,0 +1,33 @@
|
||||
(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).
|
||||
@@ -0,0 +1,9 @@
|
||||
(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
|
||||
@@ -0,0 +1,18 @@
|
||||
(hunyuanvideo)=
|
||||
|
||||
# HunyuanVideo
|
||||
## Inference HunyuanVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model
|
||||
```
|
||||
|
||||
We provide two examples in the following script to run inference with STA + [TeaCache](https://github.com/ali-vilab/TeaCache) and STA only.
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_hunyuan_STA.sh
|
||||
```
|
||||
|
||||
## Video Demos using STA + Teacache
|
||||
Visit our [demo website](https://fast-video.github.io/) to explore our complete collection of examples. We shorten a single video generation process from 945s to 317s on H100.
|
||||
@@ -0,0 +1,16 @@
|
||||
(stepvideo)=
|
||||
|
||||
# StepVideo
|
||||
## Inference StepVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
```
|
||||
python scripts/huggingface/download_hf.py --repo_id=stepfun-ai/stepvideo-t2v --local_dir=data/stepvideo-t2v --repo_type=model
|
||||
```
|
||||
|
||||
Use the following scripts to run inference for StepVideo. When using STA for inference, the generated videos will have dimensions of 204×768×768 (currently, this is the only supported shape).
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_stepvideo_STA.sh # Inference stepvideo with STA
|
||||
sh scripts/inference/inference_stepvideo.sh # Inference original stepvideo
|
||||
```
|
||||
@@ -0,0 +1,11 @@
|
||||
(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>
|
||||
@@ -0,0 +1,25 @@
|
||||
(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
|
||||
```
|
||||
@@ -0,0 +1,7 @@
|
||||
(sta-test)=
|
||||
|
||||
# Test
|
||||
|
||||
```bash
|
||||
python test/test_sta.py
|
||||
```
|
||||
@@ -0,0 +1,17 @@
|
||||
(sta-usage)=
|
||||
|
||||
# Usage
|
||||
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
|
||||
```
|
||||
@@ -1,12 +0,0 @@
|
||||
#!/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
|
||||
|
||||
pip install -r requirements-lint.txt
|
||||
|
||||
# install fastvideo
|
||||
pip install -e .
|
||||
@@ -27,8 +27,7 @@ class T5dataset(Dataset):
|
||||
self.vae_debug = vae_debug
|
||||
with open(self.json_path, "r") as f:
|
||||
train_dataset = json.load(f)
|
||||
self.train_dataset = sorted(train_dataset,
|
||||
key=lambda x: x["latent_path"])
|
||||
self.train_dataset = sorted(train_dataset, key=lambda x: x["latent_path"])
|
||||
|
||||
def __getitem__(self, idx):
|
||||
caption = self.train_dataset[idx]["caption"]
|
||||
@@ -36,17 +35,13 @@ 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:
|
||||
latents = []
|
||||
|
||||
return dict(caption=caption,
|
||||
latents=latents,
|
||||
filename=filename,
|
||||
length=length)
|
||||
return dict(caption=caption, latents=latents, filename=filename, length=length)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.train_dataset)
|
||||
@@ -60,31 +55,21 @@ 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)
|
||||
os.makedirs(os.path.join(args.output_dir, "video"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "prompt_embed"), exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "prompt_attention_mask"),
|
||||
exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "prompt_attention_mask"), exist_ok=True)
|
||||
|
||||
latents_json_path = os.path.join(args.output_dir,
|
||||
"videos2caption_temp.json")
|
||||
latents_json_path = os.path.join(args.output_dir, "videos2caption_temp.json")
|
||||
train_dataset = T5dataset(latents_json_path, args.vae_debug)
|
||||
text_encoder = load_text_encoder(args.model_type,
|
||||
args.model_path,
|
||||
device=device)
|
||||
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,
|
||||
@@ -96,26 +81,19 @@ 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 = 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)
|
||||
torch.save(prompt_attention_mask[idx], prompt_attention_mask_path)
|
||||
print(f"sample {video_name} saved")
|
||||
if args.vae_debug:
|
||||
export_to_video(video[idx], video_path, fps=fps)
|
||||
@@ -133,8 +111,7 @@ def main(args):
|
||||
if local_rank == 0:
|
||||
# os.remove(latents_json_path)
|
||||
all_json_data = [item for sublist in gathered_data for item in sublist]
|
||||
with open(os.path.join(args.output_dir, "videos2caption.json"),
|
||||
"w") as f:
|
||||
with open(os.path.join(args.output_dir, "videos2caption.json"), "w") as f:
|
||||
json.dump(all_json_data, f, indent=4)
|
||||
|
||||
|
||||
@@ -148,8 +125,7 @@ if __name__ == "__main__":
|
||||
"--dataloader_num_workers",
|
||||
type=int,
|
||||
default=1,
|
||||
help=
|
||||
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_batch_size",
|
||||
@@ -157,16 +133,13 @@ if __name__ == "__main__":
|
||||
default=1,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument("--text_encoder_name",
|
||||
type=str,
|
||||
default="google/t5-v1_1-xxl")
|
||||
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
|
||||
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help=
|
||||
"The output directory where the model predictions and checkpoints will be written.",
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
parser.add_argument("--vae_debug", action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
@@ -20,10 +20,7 @@ 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,
|
||||
@@ -31,14 +28,10 @@ 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("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)
|
||||
@@ -48,12 +41,10 @@ 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]
|
||||
@@ -67,8 +58,7 @@ def main(args):
|
||||
dist.all_gather_object(gathered_data, local_data)
|
||||
if local_rank == 0:
|
||||
all_json_data = [item for sublist in gathered_data for item in sublist]
|
||||
with open(os.path.join(args.output_dir, "videos2caption_temp.json"),
|
||||
"w") as f:
|
||||
with open(os.path.join(args.output_dir, "videos2caption_temp.json"), "w") as f:
|
||||
json.dump(all_json_data, f, indent=4)
|
||||
|
||||
|
||||
@@ -83,8 +73,7 @@ if __name__ == "__main__":
|
||||
"--dataloader_num_workers",
|
||||
type=int,
|
||||
default=1,
|
||||
help=
|
||||
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_batch_size",
|
||||
@@ -92,15 +81,10 @@ 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)
|
||||
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--dataset", default="t2v")
|
||||
@@ -110,25 +94,21 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--speed_factor", type=float, default=1.0)
|
||||
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--text_encoder_name",
|
||||
type=str,
|
||||
default="google/t5-v1_1-xxl")
|
||||
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
|
||||
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
|
||||
parser.add_argument("--cfg", type=float, default=0.0)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help=
|
||||
"The output directory where the model predictions and checkpoints will be written.",
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=
|
||||
("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
@@ -18,14 +18,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)
|
||||
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
|
||||
# output_dir/validation/prompt_attention_mask
|
||||
# output_dir/validation/prompt_embed
|
||||
@@ -34,8 +29,7 @@ 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)
|
||||
|
||||
with open(args.validation_prompt_txt, "r", encoding="utf-8") as file:
|
||||
lines = file.readlines()
|
||||
@@ -43,12 +37,9 @@ def main(args):
|
||||
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",
|
||||
@@ -56,8 +47,7 @@ def main(args):
|
||||
f"{file_name}.pt",
|
||||
)
|
||||
torch.save(prompt_embeds[0], prompt_embed_path)
|
||||
torch.save(prompt_attention_mask[0],
|
||||
prompt_attention_mask_path)
|
||||
torch.save(prompt_attention_mask[0], prompt_attention_mask_path)
|
||||
print(f"sample {file_name} saved")
|
||||
|
||||
|
||||
@@ -71,8 +61,7 @@ if __name__ == "__main__":
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help=
|
||||
"The output directory where the model predictions and checkpoints will be written.",
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
|
||||
@@ -3,16 +3,14 @@ from torchvision.transforms import Lambda
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.dataset.t2v_datasets import T2V_dataset
|
||||
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
from fastvideo.dataset.transform import CenterCropResizeVideo, Normalize255, TemporalRandomCrop
|
||||
|
||||
|
||||
def getdataset(args):
|
||||
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
|
||||
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
|
||||
resize_topcrop = [
|
||||
CenterCropResizeVideo((args.max_height, args.max_width),
|
||||
top_crop=True),
|
||||
CenterCropResizeVideo((args.max_height, args.max_width), top_crop=True),
|
||||
]
|
||||
resize = [
|
||||
CenterCropResizeVideo((args.max_height, args.max_width)),
|
||||
@@ -27,8 +25,7 @@ def getdataset(args):
|
||||
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,
|
||||
@@ -66,8 +63,7 @@ if __name__ == "__main__":
|
||||
"interpolation_scale_h": 1,
|
||||
"interpolation_scale_w": 1,
|
||||
"cache_dir": "../cache_dir",
|
||||
"image_data":
|
||||
"/storage/ongoing/new/Open-Sora-Plan-bak/7.14bak/scripts/train_data/image_data.txt",
|
||||
"image_data": "/storage/ongoing/new/Open-Sora-Plan-bak/7.14bak/scripts/train_data/image_data.txt",
|
||||
"video_data": "1",
|
||||
"train_fps": 24,
|
||||
"drop_short_ratio": 1.0,
|
||||
@@ -84,10 +80,7 @@ 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:
|
||||
|
||||
@@ -20,10 +20,8 @@ class LatentDataset(Dataset):
|
||||
self.datase_dir_path = os.path.dirname(json_path)
|
||||
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_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")
|
||||
with open(self.json_path, "r") as f:
|
||||
self.data_anno = json.load(f)
|
||||
# json.load(f) already keeps the order
|
||||
@@ -33,16 +31,12 @@ 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"]
|
||||
prompt_embed_file = self.data_anno[idx]["prompt_embed_path"]
|
||||
prompt_attention_mask_file = self.data_anno[idx][
|
||||
"prompt_attention_mask"]
|
||||
prompt_attention_mask_file = self.data_anno[idx]["prompt_attention_mask"]
|
||||
# load
|
||||
latent = torch.load(
|
||||
os.path.join(self.latent_dir, latent_file),
|
||||
@@ -60,8 +54,7 @@ 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,
|
||||
)
|
||||
@@ -111,13 +104,8 @@ 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/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)
|
||||
for latent, prompt_embed, latent_attn_mask, prompt_attention_mask in dataloader:
|
||||
print(
|
||||
latent.shape,
|
||||
|
||||
@@ -46,8 +46,7 @@ class DataSetProg(metaclass=SingletonMeta):
|
||||
|
||||
for i in range(self.num_workers):
|
||||
self.n_used_elements[i] = 0
|
||||
per_worker = int(
|
||||
math.ceil(len(self.elements) / float(self.num_workers)))
|
||||
per_worker = int(math.ceil(len(self.elements) / float(self.num_workers)))
|
||||
start = i * per_worker
|
||||
end = min(start + per_worker, len(self.elements))
|
||||
self.worker_elements[i] = self.elements[start:end]
|
||||
@@ -58,9 +57,7 @@ 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
|
||||
|
||||
@@ -68,10 +65,7 @@ class DataSetProg(metaclass=SingletonMeta):
|
||||
dataset_prog = DataSetProg()
|
||||
|
||||
|
||||
def filter_resolution(h,
|
||||
w,
|
||||
max_h_div_w_ratio=17 / 16,
|
||||
min_h_div_w_ratio=8 / 16):
|
||||
def filter_resolution(h, w, max_h_div_w_ratio=17 / 16, min_h_div_w_ratio=8 / 16):
|
||||
if h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio:
|
||||
return True
|
||||
return False
|
||||
@@ -79,8 +73,7 @@ def filter_resolution(h,
|
||||
|
||||
class T2V_dataset(Dataset):
|
||||
|
||||
def __init__(self, args, transform, temporal_sample, tokenizer,
|
||||
transform_topcrop):
|
||||
def __init__(self, args, transform, temporal_sample, tokenizer, transform_topcrop):
|
||||
self.data = args.data_merge_path
|
||||
self.num_frames = args.num_frames
|
||||
self.train_fps = args.train_fps
|
||||
@@ -109,8 +102,7 @@ class T2V_dataset(Dataset):
|
||||
self.lengths = self.sample_num_frames
|
||||
|
||||
n_elements = len(cap_list)
|
||||
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list,
|
||||
n_elements)
|
||||
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list, n_elements)
|
||||
|
||||
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
|
||||
|
||||
@@ -137,8 +129,7 @@ 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,8 +169,7 @@ class T2V_dataset(Dataset):
|
||||
)
|
||||
|
||||
def get_image(self, idx):
|
||||
image_data = dataset_prog.cap_list[
|
||||
idx] # [{'path': path, 'cap': cap}, ...]
|
||||
image_data = dataset_prog.cap_list[idx] # [{'path': path, 'cap': cap}, ...]
|
||||
|
||||
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
|
||||
image = torch.from_numpy(np.array(image)) # [h, w, c]
|
||||
@@ -188,15 +178,13 @@ 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)
|
||||
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 = [], []
|
||||
@@ -250,12 +238,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"]
|
||||
height, width = i["resolution"]["height"], i["resolution"]["width"]
|
||||
aspect = self.max_height / self.max_width
|
||||
hw_aspect_thr = 1.5
|
||||
is_pick = filter_resolution(
|
||||
@@ -273,34 +259,29 @@ 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
|
||||
|
||||
# too long video will be temporal-crop randomly
|
||||
if len(frame_indices) > self.num_frames:
|
||||
begin_index, end_index = self.temporal_sample(
|
||||
len(frame_indices))
|
||||
begin_index, end_index = self.temporal_sample(len(frame_indices))
|
||||
frame_indices = frame_indices[begin_index:end_index]
|
||||
# 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
|
||||
@@ -309,32 +290,26 @@ 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 extension {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):
|
||||
decord_vr = self.v_decoder(path)
|
||||
video_data = decord_vr.get_batch(frame_indices).asnumpy()
|
||||
video_data = torch.from_numpy(video_data)
|
||||
video_data = video_data.permute(0, 3, 1,
|
||||
2) # (T, H, W, C) -> (T C H W)
|
||||
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
|
||||
return video_data
|
||||
|
||||
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:
|
||||
|
||||
@@ -21,19 +21,15 @@ 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):
|
||||
@@ -48,9 +44,7 @@ def crop(clip, i, j, h, 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,
|
||||
@@ -62,9 +56,7 @@ 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(
|
||||
@@ -174,8 +166,7 @@ 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
|
||||
|
||||
@@ -236,9 +227,7 @@ 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
|
||||
@@ -312,9 +301,7 @@ 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:
|
||||
@@ -334,8 +321,7 @@ class CenterCropResizeVideo:
|
||||
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
|
||||
@@ -349,10 +335,7 @@ 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,
|
||||
@@ -378,9 +361,7 @@ class UCFCenterCropVideo:
|
||||
):
|
||||
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)
|
||||
@@ -395,9 +376,7 @@ 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
|
||||
|
||||
@@ -417,9 +396,7 @@ class KineticsRandomCropResizeVideo:
|
||||
):
|
||||
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)
|
||||
@@ -428,8 +405,7 @@ class KineticsRandomCropResizeVideo:
|
||||
|
||||
def __call__(self, clip):
|
||||
clip_random_crop = random_shift_crop(clip)
|
||||
clip_resize = resize(clip_random_crop, self.size,
|
||||
self.interpolation_mode)
|
||||
clip_resize = resize(clip_random_crop, self.size, self.interpolation_mode)
|
||||
return clip_resize
|
||||
|
||||
|
||||
@@ -442,9 +418,7 @@ class CenterCropVideo:
|
||||
):
|
||||
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)
|
||||
@@ -571,8 +545,7 @@ 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
|
||||
@@ -587,18 +560,14 @@ if __name__ == "__main__":
|
||||
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),
|
||||
transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5], inplace=True),
|
||||
])
|
||||
|
||||
target_video_len = 32
|
||||
@@ -613,10 +582,7 @@ 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]
|
||||
@@ -627,14 +593,11 @@ 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)
|
||||
|
||||
io.write_video("./test.avi",
|
||||
select_vframes_trans_int.permute(0, 2, 3, 1),
|
||||
fps=8)
|
||||
io.write_video("./test.avi", select_vframes_trans_int.permute(0, 2, 3, 1), fps=8)
|
||||
|
||||
for i in range(target_video_len):
|
||||
save_image(
|
||||
|
||||
+100
-215
@@ -21,23 +21,17 @@ 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.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.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.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.parallel_states import (destroy_sequence_parallel_group, get_sequence_parallel_state,
|
||||
initialize_sequence_parallel_state)
|
||||
from fastvideo.utils.validation import log_validation
|
||||
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
@@ -59,18 +53,14 @@ 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)
|
||||
absolute_mean = torch.mean(
|
||||
torch.abs(model_pred)) / gradient_accumulation_steps
|
||||
absolute_max = torch.max(
|
||||
torch.abs(model_pred)) / 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["largest singular value"] += torch.mean(
|
||||
largest_singular_value).item()
|
||||
norms["largest singular value"] += torch.mean(largest_singular_value).item()
|
||||
norms["absolute mean"] += absolute_mean.item()
|
||||
norms["absolute max"] += absolute_max.item()
|
||||
|
||||
@@ -118,23 +108,18 @@ 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.
|
||||
# sigmas = get_sigmas(start_timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
|
||||
sigmas = extract_into_tensor(solver.sigmas, index, model_input.shape)
|
||||
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index,
|
||||
model_input.shape)
|
||||
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index, model_input.shape)
|
||||
|
||||
timesteps = (sigmas *
|
||||
noise_scheduler.config.num_train_timesteps).view(-1)
|
||||
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):
|
||||
@@ -146,15 +131,13 @@ 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):
|
||||
@@ -177,10 +160,8 @@ def distill_one_step(
|
||||
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)
|
||||
x_prev = solver.euler_step(noisy_model_input, teacher_output,
|
||||
index)
|
||||
teacher_output = uncond_teacher_output + w * (cond_teacher_output - uncond_teacher_output)
|
||||
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
|
||||
|
||||
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
|
||||
with torch.no_grad():
|
||||
@@ -202,33 +183,26 @@ 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")
|
||||
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()
|
||||
@@ -238,12 +212,9 @@ 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()
|
||||
@@ -306,11 +277,8 @@ 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,
|
||||
@@ -322,12 +290,9 @@ def main(args):
|
||||
if args.use_lora:
|
||||
transformer.config.lora_rank = args.lora_rank
|
||||
transformer.config.lora_alpha = args.lora_alpha
|
||||
transformer.config.lora_target_modules = [
|
||||
"to_k", "to_q", "to_v", "to_out.0"
|
||||
]
|
||||
transformer.config.lora_target_modules = ["to_k", "to_q", "to_v", "to_out.0"]
|
||||
transformer._no_split_modules = no_split_modules
|
||||
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](
|
||||
transformer)
|
||||
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
|
||||
|
||||
transformer = FSDP(
|
||||
transformer,
|
||||
@@ -345,13 +310,10 @@ def main(args):
|
||||
main_print("--> 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)
|
||||
@@ -359,8 +321,7 @@ 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,
|
||||
@@ -376,8 +337,7 @@ def main(args):
|
||||
)
|
||||
solver.to(device)
|
||||
params_to_optimize = transformer.parameters()
|
||||
params_to_optimize = list(
|
||||
filter(lambda p: p.requires_grad, params_to_optimize))
|
||||
params_to_optimize = list(filter(lambda p: p.requires_grad, params_to_optimize))
|
||||
|
||||
optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
@@ -389,8 +349,8 @@ 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
|
||||
@@ -404,8 +364,7 @@ def main(args):
|
||||
last_epoch=init_steps - 1,
|
||||
)
|
||||
|
||||
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t,
|
||||
args.cfg)
|
||||
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(
|
||||
@@ -429,42 +388,33 @@ 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)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps /
|
||||
num_update_steps_per_epoch)
|
||||
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:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
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 = (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" Gradient Accumulation steps = {args.gradient_accumulation_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" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
|
||||
main_print(f" Total optimization steps = {args.max_train_steps}")
|
||||
main_print(
|
||||
f" Total training parameters per FSDP shard = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e9} B"
|
||||
)
|
||||
# print dtype
|
||||
main_print(
|
||||
f" Master weight dtype: {transformer.parameters().__next__().dtype}")
|
||||
main_print(f" Master weight dtype: {transformer.parameters().__next__().dtype}")
|
||||
|
||||
# Potentially load in the weights and states from a previous save
|
||||
if args.resume_from_checkpoint:
|
||||
assert NotImplementedError(
|
||||
"resume_from_checkpoint is not supported now.")
|
||||
assert NotImplementedError("resume_from_checkpoint is not supported now.")
|
||||
# TODO
|
||||
|
||||
progress_bar = tqdm(
|
||||
@@ -546,37 +496,26 @@ def main(args):
|
||||
if rank <= 0:
|
||||
wandb.log(
|
||||
{
|
||||
"train_loss":
|
||||
loss,
|
||||
"learning_rate":
|
||||
lr_scheduler.get_last_lr()[0],
|
||||
"step_time":
|
||||
step_time,
|
||||
"avg_step_time":
|
||||
avg_step_time,
|
||||
"grad_norm":
|
||||
grad_norm,
|
||||
"pred_fro_norm":
|
||||
pred_norm["fro"], # codespell:ignore
|
||||
"pred_largest_singular_value":
|
||||
pred_norm["largest singular value"],
|
||||
"pred_absolute_mean":
|
||||
pred_norm["absolute mean"],
|
||||
"pred_absolute_max":
|
||||
pred_norm["absolute max"],
|
||||
"train_loss": loss,
|
||||
"learning_rate": lr_scheduler.get_last_lr()[0],
|
||||
"step_time": step_time,
|
||||
"avg_step_time": avg_step_time,
|
||||
"grad_norm": grad_norm,
|
||||
"pred_fro_norm": pred_norm["fro"], # codespell:ignore
|
||||
"pred_largest_singular_value": pred_norm["largest singular value"],
|
||||
"pred_absolute_mean": pred_norm["absolute mean"],
|
||||
"pred_absolute_max": pred_norm["absolute max"],
|
||||
},
|
||||
step=step,
|
||||
)
|
||||
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:
|
||||
save_checkpoint(ema_transformer, rank, args.output_dir,
|
||||
step)
|
||||
save_checkpoint(ema_transformer, rank, args.output_dir, step)
|
||||
else:
|
||||
save_checkpoint(transformer, rank, args.output_dir, step)
|
||||
dist.barrier()
|
||||
@@ -610,11 +549,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)
|
||||
save_checkpoint(transformer, rank, args.output_dir, args.max_train_steps)
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
destroy_sequence_parallel_group()
|
||||
@@ -623,10 +560,7 @@ 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)
|
||||
@@ -637,8 +571,7 @@ if __name__ == "__main__":
|
||||
"--dataloader_num_workers",
|
||||
type=int,
|
||||
default=10,
|
||||
help=
|
||||
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_batch_size",
|
||||
@@ -646,10 +579,7 @@ 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
|
||||
|
||||
@@ -671,16 +601,12 @@ 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,
|
||||
default=None,
|
||||
help=
|
||||
"The output directory where the model predictions and checkpoints will be written.",
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--checkpoints_total_limit",
|
||||
@@ -692,37 +618,31 @@ 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
|
||||
@@ -731,29 +651,25 @@ if __name__ == "__main__":
|
||||
"--max_train_steps",
|
||||
type=int,
|
||||
default=None,
|
||||
help=
|
||||
"Total number of training steps to perform. If provided, overrides num_train_epochs.",
|
||||
help="Total number of training steps to perform. If provided, overrides num_train_epochs.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--gradient_accumulation_steps",
|
||||
type=int,
|
||||
default=1,
|
||||
help=
|
||||
"Number of updates steps to accumulate before performing a backward/update pass.",
|
||||
help="Number of updates steps to accumulate before performing a backward/update pass.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--learning_rate",
|
||||
type=float,
|
||||
default=1e-4,
|
||||
help=
|
||||
"Initial learning rate (after the potential warmup period) to use.",
|
||||
help="Initial learning rate (after the potential warmup period) to use.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--scale_lr",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help=
|
||||
"Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
|
||||
help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lr_warmup_steps",
|
||||
@@ -761,47 +677,36 @@ 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",
|
||||
help=
|
||||
"Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
|
||||
help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
|
||||
)
|
||||
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
|
||||
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",
|
||||
type=str,
|
||||
default=None,
|
||||
choices=["no", "fp16", "bf16"],
|
||||
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."
|
||||
),
|
||||
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."),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use_cpu_offload",
|
||||
action="store_true",
|
||||
help=
|
||||
"Whether to use CPU offload for param & gradient & optimizer states.",
|
||||
help="Whether to use CPU offload for param & gradient & optimizer states.",
|
||||
)
|
||||
|
||||
parser.add_argument("--sp_size",
|
||||
type=int,
|
||||
default=1,
|
||||
help="For sequence parallel")
|
||||
parser.add_argument("--sp_size", type=int, default=1, help="For sequence parallel")
|
||||
parser.add_argument(
|
||||
"--train_sp_batch_size",
|
||||
type=int,
|
||||
@@ -815,14 +720,8 @@ 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
|
||||
@@ -830,9 +729,8 @@ 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(
|
||||
@@ -852,15 +750,9 @@ 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,
|
||||
@@ -873,16 +765,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("--use_ema",
|
||||
action="store_true",
|
||||
help="Whether to use EMA.")
|
||||
parser.add_argument("--multi_phased_distill_schedule",
|
||||
type=str,
|
||||
default=None)
|
||||
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)
|
||||
parser.add_argument("--pred_decay_type", default="l1")
|
||||
parser.add_argument("--hunyuan_teacher_disable_cfg", action="store_true")
|
||||
|
||||
@@ -12,16 +12,12 @@ class DiscriminatorHead(nn.Module):
|
||||
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)
|
||||
@@ -53,10 +49,8 @@ class Discriminator(nn.Module):
|
||||
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
|
||||
nn.ModuleList([DiscriminatorHead(adapter_channel) for _ in range(self.num_h_per_head)])
|
||||
for adapter_channel in adapter_channel_dims
|
||||
])
|
||||
|
||||
def forward(self, features):
|
||||
|
||||
+24
-54
@@ -39,28 +39,21 @@ 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
|
||||
(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
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
self.sigmas = self.sigmas.to(
|
||||
"cpu") # to avoid too much CPU/GPU communication
|
||||
self.sigmas = self.sigmas.to("cpu") # to avoid too much CPU/GPU communication
|
||||
self.sigma_min = self.sigmas[-1].item()
|
||||
self.sigma_max = self.sigmas[0].item()
|
||||
|
||||
@@ -119,9 +112,7 @@ 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).
|
||||
|
||||
@@ -132,19 +123,14 @@ 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
|
||||
|
||||
@@ -208,10 +194,9 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
|
||||
|
||||
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."), )
|
||||
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)
|
||||
@@ -241,18 +226,14 @@ 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_prev = np.asarray(
|
||||
[0] + self.euler_timesteps[:-1].tolist())
|
||||
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()
|
||||
self.euler_timesteps_prev = torch.from_numpy(self.euler_timesteps_prev).long()
|
||||
self.sigmas = torch.from_numpy(self.sigmas)
|
||||
self.sigmas_prev = torch.from_numpy(self.sigmas_prev)
|
||||
|
||||
@@ -265,10 +246,8 @@ class EulerSolver:
|
||||
return self
|
||||
|
||||
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 = extract_into_tensor(self.sigmas, 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
|
||||
|
||||
@@ -280,29 +259,20 @@ class EulerSolver:
|
||||
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 = valid_indices_mask.flip(dims=[1]).long().argmax(dim=1)
|
||||
last_valid_index = inference_indices.size(0) - 1 - last_valid_index
|
||||
timestep_index_end = inference_indices[last_valid_index]
|
||||
|
||||
if is_target:
|
||||
sigma = extract_into_tensor(self.sigmas_prev, timestep_index,
|
||||
sample.shape)
|
||||
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 = extract_into_tensor(self.sigmas, timestep_index, 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
|
||||
|
||||
+81
-178
@@ -20,27 +20,20 @@ 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.dataset.latent_datasets import (LatentDataset, latent_collate_function)
|
||||
from fastvideo.distill.discriminator import Discriminator
|
||||
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
|
||||
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.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.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.fsdp_util import (apply_fsdp_checkpointing, get_discriminator_fsdp_kwargs, get_dit_fsdp_kwargs)
|
||||
from fastvideo.utils.load import load_transformer
|
||||
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.parallel_states import (destroy_sequence_parallel_group, get_sequence_parallel_state,
|
||||
initialize_sequence_parallel_state)
|
||||
from fastvideo.utils.validation import log_validation
|
||||
|
||||
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
|
||||
@@ -83,9 +76,8 @@ 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
|
||||
|
||||
|
||||
@@ -111,8 +103,8 @@ def gan_g_loss(
|
||||
)[1]
|
||||
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
|
||||
|
||||
|
||||
@@ -151,22 +143,18 @@ 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.
|
||||
# sigmas = get_sigmas(start_timesteps, n_dim=model_input.ndim, dtype=model_input.dtype)
|
||||
sigmas = extract_into_tensor(solver.sigmas, index, model_input.shape)
|
||||
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index,
|
||||
model_input.shape)
|
||||
sigmas_prev = extract_into_tensor(solver.sigmas_prev, index, model_input.shape)
|
||||
|
||||
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
|
||||
@@ -180,8 +168,7 @@ 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)
|
||||
|
||||
# # simplified flow matching aka 0-rectified flow matching loss
|
||||
# # target = model_input - noise
|
||||
@@ -196,12 +183,9 @@ def distill_one_step_adv(
|
||||
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_adv = (sigmas_adv *
|
||||
noise_scheduler.config.num_train_timesteps).view(-1)
|
||||
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_adv = (sigmas_adv * noise_scheduler.config.num_train_timesteps).view(-1)
|
||||
|
||||
with torch.no_grad():
|
||||
w = distill_cfg
|
||||
@@ -225,8 +209,7 @@ 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
|
||||
@@ -240,20 +223,14 @@ 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)
|
||||
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)
|
||||
(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(
|
||||
@@ -358,9 +335,7 @@ 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,
|
||||
@@ -368,18 +343,14 @@ def main(args):
|
||||
args.use_cpu_offload,
|
||||
args.master_weight_type,
|
||||
)
|
||||
discriminator_fsdp_kwargs = get_discriminator_fsdp_kwargs(
|
||||
args.master_weight_type)
|
||||
discriminator_fsdp_kwargs = get_discriminator_fsdp_kwargs(args.master_weight_type)
|
||||
if args.use_lora:
|
||||
assert args.model_type == "mochi", "LoRA is only supported for Mochi model."
|
||||
transformer.config.lora_rank = args.lora_rank
|
||||
transformer.config.lora_alpha = args.lora_alpha
|
||||
transformer.config.lora_target_modules = [
|
||||
"to_k", "to_q", "to_v", "to_out.0"
|
||||
]
|
||||
transformer.config.lora_target_modules = ["to_k", "to_q", "to_v", "to_out.0"]
|
||||
transformer._no_split_modules = no_split_modules
|
||||
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](
|
||||
transformer)
|
||||
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
|
||||
|
||||
transformer = FSDP(
|
||||
transformer,
|
||||
@@ -396,18 +367,14 @@ def main(args):
|
||||
main_print("--> 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
|
||||
@@ -418,8 +385,7 @@ def main(args):
|
||||
)
|
||||
solver.to(device)
|
||||
params_to_optimize = transformer.parameters()
|
||||
params_to_optimize = list(
|
||||
filter(lambda p: p.requires_grad, params_to_optimize))
|
||||
params_to_optimize = list(filter(lambda p: p.requires_grad, params_to_optimize))
|
||||
|
||||
optimizer = torch.optim.AdamW(
|
||||
params_to_optimize,
|
||||
@@ -439,8 +405,8 @@ 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)
|
||||
elif args.resume_from_checkpoint:
|
||||
(
|
||||
transformer,
|
||||
@@ -469,8 +435,7 @@ def main(args):
|
||||
last_epoch=init_steps - 1,
|
||||
)
|
||||
|
||||
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t,
|
||||
args.cfg)
|
||||
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(
|
||||
@@ -494,37 +459,29 @@ 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)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps /
|
||||
num_update_steps_per_epoch)
|
||||
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:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
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 = (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" Gradient Accumulation steps = {args.gradient_accumulation_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" Gradient Accumulation steps = {args.gradient_accumulation_steps}")
|
||||
main_print(f" Total optimization steps = {args.max_train_steps}")
|
||||
main_print(
|
||||
f" Total training parameters per FSDP shard = {sum(p.numel() for p in transformer.parameters() if p.requires_grad) / 1e9} B"
|
||||
)
|
||||
# print dtype
|
||||
main_print(
|
||||
f" Master weight dtype: {transformer.parameters().__next__().dtype}")
|
||||
main_print(f" Master weight dtype: {transformer.parameters().__next__().dtype}")
|
||||
|
||||
progress_bar = tqdm(
|
||||
range(0, args.max_train_steps),
|
||||
@@ -619,8 +576,7 @@ 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
|
||||
@@ -652,11 +608,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)
|
||||
save_checkpoint(transformer, rank, args.output_dir, args.max_train_steps)
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
destroy_sequence_parallel_group()
|
||||
@@ -665,10 +619,7 @@ 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)
|
||||
@@ -678,8 +629,7 @@ if __name__ == "__main__":
|
||||
"--dataloader_num_workers",
|
||||
type=int,
|
||||
default=10,
|
||||
help=
|
||||
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--train_batch_size",
|
||||
@@ -687,10 +637,7 @@ 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
|
||||
|
||||
@@ -709,16 +656,12 @@ 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,
|
||||
default=None,
|
||||
help=
|
||||
"The output directory where the model predictions and checkpoints will be written.",
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--checkpoints_total_limit",
|
||||
@@ -730,10 +673,9 @@ 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)
|
||||
@@ -741,27 +683,22 @@ 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
|
||||
@@ -770,29 +707,25 @@ if __name__ == "__main__":
|
||||
"--max_train_steps",
|
||||
type=int,
|
||||
default=None,
|
||||
help=
|
||||
"Total number of training steps to perform. If provided, overrides num_train_epochs.",
|
||||
help="Total number of training steps to perform. If provided, overrides num_train_epochs.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--learning_rate",
|
||||
type=float,
|
||||
default=1e-4,
|
||||
help=
|
||||
"Initial learning rate (after the potential warmup period) to use.",
|
||||
help="Initial learning rate (after the potential warmup period) to use.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--discriminator_learning_rate",
|
||||
type=float,
|
||||
default=1e-5,
|
||||
help=
|
||||
"Initial learning rate (after the potential warmup period) to use.",
|
||||
help="Initial learning rate (after the potential warmup period) to use.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--scale_lr",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help=
|
||||
"Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
|
||||
help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--lr_warmup_steps",
|
||||
@@ -800,47 +733,36 @@ 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",
|
||||
help=
|
||||
"Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
|
||||
help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.",
|
||||
)
|
||||
parser.add_argument("--selective_checkpointing", type=float, default=1.0)
|
||||
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",
|
||||
type=str,
|
||||
default=None,
|
||||
choices=["no", "fp16", "bf16"],
|
||||
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."
|
||||
),
|
||||
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."),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use_cpu_offload",
|
||||
action="store_true",
|
||||
help=
|
||||
"Whether to use CPU offload for param & gradient & optimizer states.",
|
||||
help="Whether to use CPU offload for param & gradient & optimizer states.",
|
||||
)
|
||||
|
||||
parser.add_argument("--sp_size",
|
||||
type=int,
|
||||
default=1,
|
||||
help="For sequence parallel")
|
||||
parser.add_argument("--sp_size", type=int, default=1, help="For sequence parallel")
|
||||
parser.add_argument(
|
||||
"--train_sp_batch_size",
|
||||
type=int,
|
||||
@@ -854,24 +776,15 @@ 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("--multi_phased_distill_schedule", type=str, default=None)
|
||||
parser.add_argument(
|
||||
"--gradient_accumulation_steps",
|
||||
type=int,
|
||||
default=1,
|
||||
help=
|
||||
"Number of updates steps to accumulate before performing a backward/update pass.",
|
||||
help="Number of updates steps to accumulate before performing a backward/update pass.",
|
||||
)
|
||||
|
||||
# lr_scheduler
|
||||
@@ -879,9 +792,8 @@ 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(
|
||||
@@ -901,15 +813,9 @@ 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,
|
||||
@@ -928,10 +834,7 @@ 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(
|
||||
"--linear_quadratic_threshold",
|
||||
type=float,
|
||||
|
||||
@@ -1,48 +1,27 @@
|
||||
try:
|
||||
from flash_attn_interface import flash_attn_func as flash_attn_func_v3
|
||||
has_v3 = True
|
||||
except:
|
||||
pass
|
||||
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, attn_impl="fa3",
|
||||
):
|
||||
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)
|
||||
if attn_impl == "fa3":
|
||||
assert has_v3
|
||||
q, k, v = x_unpad[:, 0].unsqueeze(0), x_unpad[:, 1].unsqueeze(0), x_unpad[:, 2].unsqueeze(0)
|
||||
assert dropout_p == 0.0
|
||||
output_unpad = flash_attn_func_v3(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
)[0].squeeze(0)
|
||||
else:
|
||||
output_unpad = flash_attn_varlen_qkvpacked_func(
|
||||
x_unpad,
|
||||
cu_seqlens,
|
||||
max_s,
|
||||
dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=causal,
|
||||
)
|
||||
output_unpad = flash_attn_varlen_qkvpacked_func(
|
||||
x_unpad,
|
||||
cu_seqlens,
|
||||
max_s,
|
||||
dropout_p,
|
||||
softmax_scale=softmax_scale,
|
||||
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,
|
||||
)
|
||||
|
||||
@@ -32,14 +32,13 @@ 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, BaseOutput, deprecate, logging, replace_example_docstring,
|
||||
scale_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 fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
|
||||
from ...constants import PRECISION_TO_TYPE
|
||||
from ...modules import HYVideoDiffusionTransformer
|
||||
@@ -56,14 +55,12 @@ 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
|
||||
|
||||
|
||||
@@ -99,28 +96,22 @@ 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)
|
||||
@@ -158,9 +149,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
model_cpu_offload_seq = "text_encoder->text_encoder_2->transformer->vae"
|
||||
_optional_components = ["text_encoder_2"]
|
||||
_exclude_from_cpu_offload = ["transformer"]
|
||||
_callback_tensor_inputs = [
|
||||
"latents", "prompt_embeds", "negative_prompt_embeds"
|
||||
]
|
||||
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -184,8 +173,7 @@ 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 "
|
||||
@@ -193,27 +181,19 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
" 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)
|
||||
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)
|
||||
@@ -225,10 +205,8 @@ 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.image_processor = VaeImageProcessor(
|
||||
vae_scale_factor=self.vae_scale_factor)
|
||||
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(
|
||||
self,
|
||||
@@ -296,14 +274,11 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
if prompt_embeds is None:
|
||||
# textual inversion: process multi-vector tokens if necessary
|
||||
if isinstance(self, TextualInversionLoaderMixin):
|
||||
prompt = self.maybe_convert_prompt(prompt,
|
||||
text_encoder.tokenizer)
|
||||
prompt = self.maybe_convert_prompt(prompt, text_encoder.tokenizer)
|
||||
|
||||
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(
|
||||
@@ -315,23 +290,19 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
# Access the `hidden_states` first, that contains a tuple of
|
||||
# all the hidden states from the encoder layers. Then index into
|
||||
# the tuple to access the hidden states from the desired layer.
|
||||
prompt_embeds = prompt_outputs.hidden_states_list[-(clip_skip +
|
||||
1)]
|
||||
prompt_embeds = prompt_outputs.hidden_states_list[-(clip_skip + 1)]
|
||||
# We also need to apply the final LayerNorm here to not mess with the
|
||||
# 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.repeat(1, num_videos_per_prompt)
|
||||
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
|
||||
@@ -340,21 +311,18 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
else:
|
||||
prompt_embeds_dtype = prompt_embeds.dtype
|
||||
|
||||
prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype,
|
||||
device=device)
|
||||
prompt_embeds = prompt_embeds.to(dtype=prompt_embeds_dtype, device=device)
|
||||
|
||||
if prompt_embeds.ndim == 2:
|
||||
bs_embed, _ = 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)
|
||||
prompt_embeds = prompt_embeds.view(
|
||||
bs_embed * num_videos_per_prompt, -1)
|
||||
prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, -1)
|
||||
else:
|
||||
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,
|
||||
@@ -365,10 +333,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
|
||||
def decode_latents(self, latents, enable_tiling=True):
|
||||
deprecation_message = "The decode_latents method is deprecated and will be removed in 1.0.0. Please use VaeImageProcessor.postprocess(...) instead"
|
||||
deprecate("decode_latents",
|
||||
"1.0.0",
|
||||
deprecation_message,
|
||||
standard_warn=False)
|
||||
deprecate("decode_latents", "1.0.0", deprecation_message, standard_warn=False)
|
||||
|
||||
latents = 1 / self.vae.config.scaling_factor * latents
|
||||
if enable_tiling:
|
||||
@@ -409,30 +374,21 @@ 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]}"
|
||||
)
|
||||
@@ -443,19 +399,13 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
" 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:
|
||||
@@ -486,14 +436,10 @@ 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)
|
||||
|
||||
@@ -585,8 +531,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
eta: float = 0.0,
|
||||
generator: Optional[Union[torch.Generator,
|
||||
List[torch.Generator]]] = None,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
@@ -597,8 +542,7 @@ 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,
|
||||
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",
|
||||
@@ -606,7 +550,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
enable_vae_sp: bool = False,
|
||||
n_tokens: Optional[int] = None,
|
||||
embedded_guidance_scale: Optional[float] = None,
|
||||
only_save_latent: bool = False,
|
||||
mask_strategy: Optional[Dict[str, list]] = None,
|
||||
**kwargs,
|
||||
):
|
||||
r"""
|
||||
@@ -707,8 +651,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
"Passing `callback_steps` as an input argument to `__call__` is deprecated, consider using `callback_on_step_end`",
|
||||
)
|
||||
|
||||
if isinstance(callback_on_step_end,
|
||||
(PipelineCallback, MultiPipelineCallbacks)):
|
||||
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
|
||||
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
|
||||
|
||||
# 0. Default height and width to unet
|
||||
@@ -744,8 +687,7 @@ 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)
|
||||
@@ -805,15 +747,13 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
if prompt_mask is not None:
|
||||
prompt_mask = torch.cat([negative_prompt_mask, prompt_mask])
|
||||
if prompt_embeds_2 is not None:
|
||||
prompt_embeds_2 = torch.cat(
|
||||
[negative_prompt_embeds_2, prompt_embeds_2])
|
||||
prompt_embeds_2 = torch.cat([negative_prompt_embeds_2, prompt_embeds_2])
|
||||
if prompt_mask_2 is not None:
|
||||
prompt_mask_2 = torch.cat(
|
||||
[negative_prompt_mask_2, prompt_mask_2])
|
||||
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,
|
||||
@@ -842,12 +782,9 @@ 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
|
||||
@@ -860,17 +797,24 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
)
|
||||
|
||||
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
|
||||
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):
|
||||
@@ -878,38 +822,31 @@ 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)
|
||||
).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):
|
||||
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
|
||||
# concat prompt_embeds_2 and prompt_embeds. Mismatch 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]),
|
||||
(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, # [2, 16, 33, 24, 42]
|
||||
latent_model_input,
|
||||
encoder_hidden_states,
|
||||
t_expand, # [2]
|
||||
prompt_mask, # [2, 256]fpdb
|
||||
t_expand,
|
||||
prompt_mask,
|
||||
mask_strategy=mask_strategy[i],
|
||||
guidance=guidance_expand,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
@@ -917,8 +854,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
|
||||
# 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
|
||||
@@ -929,29 +865,20 @@ 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 = {}
|
||||
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)
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
prompt_embeds = callback_outputs.pop(
|
||||
"prompt_embeds", prompt_embeds)
|
||||
negative_prompt_embeds = callback_outputs.pop(
|
||||
"negative_prompt_embeds", negative_prompt_embeds)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or (
|
||||
(i + 1) > num_warmup_steps and
|
||||
(i + 1) % self.scheduler.order == 0):
|
||||
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:
|
||||
@@ -971,34 +898,23 @@ 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 # 1, 16, 32, 90, 160
|
||||
|
||||
if only_save_latent:
|
||||
latents = latents.cpu().float()
|
||||
self.maybe_free_model_hooks()
|
||||
return latents
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
latents = latents / self.vae.config.scaling_factor
|
||||
|
||||
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]
|
||||
|
||||
if expand_temporal_dim or image.shape[2] == 1:
|
||||
image = image.squeeze(2)
|
||||
|
||||
else:
|
||||
image = latents
|
||||
|
||||
|
||||
@@ -80,17 +80,14 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
|
||||
self.sigmas = sigmas
|
||||
# the value fed to model
|
||||
self.timesteps = (sigmas[:-1] *
|
||||
num_train_timesteps).to(dtype=torch.float32)
|
||||
self.timesteps = (sigmas[:-1] * num_train_timesteps).to(dtype=torch.float32)
|
||||
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
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):
|
||||
@@ -146,8 +143,7 @@ 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
|
||||
@@ -174,9 +170,7 @@ 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):
|
||||
@@ -216,10 +210,9 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
|
||||
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."), )
|
||||
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)
|
||||
@@ -232,9 +225,7 @@ 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
|
||||
|
||||
@@ -7,8 +7,7 @@ from .modules.models import HUNYUAN_VIDEO_CONFIG
|
||||
|
||||
|
||||
def parse_args(namespace=None):
|
||||
parser = argparse.ArgumentParser(
|
||||
description="HunyuanVideo inference script")
|
||||
parser = argparse.ArgumentParser(description="HunyuanVideo inference script")
|
||||
|
||||
parser = add_network_args(parser)
|
||||
parser = add_extra_models_args(parser)
|
||||
@@ -36,8 +35,7 @@ def add_network_args(parser: argparse.ArgumentParser):
|
||||
"--latent-channels",
|
||||
type=str,
|
||||
default=16,
|
||||
help=
|
||||
"Number of latent channels of DiT. If None, it will be determined by `vae`. If provided, "
|
||||
help="Number of latent channels of DiT. If None, it will be determined by `vae`. If provided, "
|
||||
"it still needs to match the latent channels of the VAE model.",
|
||||
)
|
||||
group.add_argument(
|
||||
@@ -45,22 +43,16 @@ def add_network_args(parser: argparse.ArgumentParser):
|
||||
type=str,
|
||||
default="bf16",
|
||||
choices=PRECISIONS,
|
||||
help=
|
||||
"Precision mode. Options: fp32, fp16, bf16. Applied to the backbone model and optimizer.",
|
||||
help="Precision mode. Options: fp32, fp16, bf16. Applied to the backbone model and optimizer.",
|
||||
)
|
||||
|
||||
# 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(
|
||||
@@ -104,10 +96,7 @@ 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,
|
||||
@@ -138,8 +127,7 @@ def add_extra_models_args(parser: argparse.ArgumentParser):
|
||||
group.add_argument(
|
||||
"--apply-final-norm",
|
||||
action="store_true",
|
||||
help=
|
||||
"Apply final normalization to the used text encoder hidden states.",
|
||||
help="Apply final normalization to the used text encoder hidden states.",
|
||||
)
|
||||
|
||||
# - CLIP
|
||||
@@ -232,16 +220,13 @@ def add_inference_args(parser: argparse.ArgumentParser):
|
||||
"--model-base",
|
||||
type=str,
|
||||
default="ckpts",
|
||||
help=
|
||||
"Root path of all the models, including t2v models and extra models.",
|
||||
help="Root path of all the models, including t2v models and extra models.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--dit-weight",
|
||||
type=str,
|
||||
default=
|
||||
"ckpts/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
|
||||
help=
|
||||
"Path to the HunyuanVideo model. If None, search the model in the args.model_root."
|
||||
default="ckpts/hunyuan-video-t2v-720p/transformers/mp_rank_00_model_states.pt",
|
||||
help="Path to the HunyuanVideo model. If None, search the model in the args.model_root."
|
||||
"1. If it is a file, load the model directly."
|
||||
"2. If it is a directory, search the model in the directory. Support two types of models: "
|
||||
"1) named `pytorch_model_*.pt`"
|
||||
@@ -252,15 +237,13 @@ def add_inference_args(parser: argparse.ArgumentParser):
|
||||
type=str,
|
||||
default="540p",
|
||||
choices=["540p", "720p"],
|
||||
help=
|
||||
"Root path of all the models, including t2v models and extra models.",
|
||||
help="Root path of all the models, including t2v models and extra models.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--load-key",
|
||||
type=str,
|
||||
default="module",
|
||||
help=
|
||||
"Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
|
||||
help="Key to load the model states. 'module' for the main model, 'ema' for the EMA model.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--use-cpu-offload",
|
||||
@@ -284,8 +267,7 @@ def add_inference_args(parser: argparse.ArgumentParser):
|
||||
group.add_argument(
|
||||
"--disable-autocast",
|
||||
action="store_true",
|
||||
help=
|
||||
"Disable autocast for denoising loop and vae decoding in pipeline sampling.",
|
||||
help="Disable autocast for denoising loop and vae decoding in pipeline sampling.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--save-path",
|
||||
@@ -317,8 +299,7 @@ def add_inference_args(parser: argparse.ArgumentParser):
|
||||
type=int,
|
||||
nargs="+",
|
||||
default=(720, 1280),
|
||||
help=
|
||||
"Video size for training. If a single value is provided, it will be used for both height "
|
||||
help="Video size for training. If a single value is provided, it will be used for both height "
|
||||
"and width. If two values are provided, they will be used for height and width "
|
||||
"respectively.",
|
||||
)
|
||||
@@ -326,8 +307,7 @@ def add_inference_args(parser: argparse.ArgumentParser):
|
||||
"--video-length",
|
||||
type=int,
|
||||
default=129,
|
||||
help=
|
||||
"How many frames to sample from a video. if using 3d vae, the number should be 4n+1",
|
||||
help="How many frames to sample from a video. if using 3d vae, the number should be 4n+1",
|
||||
)
|
||||
# --- prompt ---
|
||||
group.add_argument(
|
||||
@@ -341,26 +321,16 @@ def add_inference_args(parser: argparse.ArgumentParser):
|
||||
type=str,
|
||||
default="auto",
|
||||
choices=["file", "random", "fixed", "auto"],
|
||||
help=
|
||||
"Seed type for evaluation. If file, use the seed from the CSV file. If random, generate a "
|
||||
help="Seed type for evaluation. If file, use the seed from the CSV file. If random, generate a "
|
||||
"random seed. If fixed, use the fixed seed given by `--seed`. If auto, `csv` will use the "
|
||||
"seed column if available, otherwise use the fixed `seed` value. `prompt` will use the "
|
||||
"fixed `seed` value.",
|
||||
)
|
||||
group.add_argument("--seed",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Seed for evaluation.")
|
||||
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,
|
||||
@@ -371,8 +341,7 @@ def add_inference_args(parser: argparse.ArgumentParser):
|
||||
group.add_argument(
|
||||
"--reproduce",
|
||||
action="store_true",
|
||||
help=
|
||||
"Enable reproducibility by setting random seeds and deterministic algorithms.",
|
||||
help="Enable reproducibility by setting random seeds and deterministic algorithms.",
|
||||
)
|
||||
|
||||
return parser
|
||||
@@ -402,14 +371,10 @@ 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
|
||||
|
||||
@@ -7,12 +7,9 @@ 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.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
|
||||
from fastvideo.models.hunyuan.diffusion.schedulers import FlowMatchDiscreteScheduler
|
||||
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
|
||||
@@ -47,17 +44,12 @@ 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
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls,
|
||||
pretrained_model_path,
|
||||
args,
|
||||
device=None,
|
||||
**kwargs):
|
||||
def from_pretrained(cls, pretrained_model_path, args, device=None, **kwargs):
|
||||
"""
|
||||
Initialize the Inference pipeline.
|
||||
|
||||
@@ -67,8 +59,7 @@ class Inference(object):
|
||||
device (int): The device for inference. Default is 0.
|
||||
"""
|
||||
# ========================================================================
|
||||
logger.info(
|
||||
f"Got text-to-video model root path: {pretrained_model_path}")
|
||||
logger.info(f"Got text-to-video model root path: {pretrained_model_path}")
|
||||
|
||||
# ==================== Initialize Distributed Environment ================
|
||||
if nccl_info.sp_size > 1:
|
||||
@@ -85,10 +76,7 @@ class Inference(object):
|
||||
|
||||
# =========================== Build main model ===========================
|
||||
logger.info("Building model...")
|
||||
factor_kwargs = {
|
||||
"device": device,
|
||||
"dtype": PRECISION_TO_TYPE[args.precision]
|
||||
}
|
||||
factor_kwargs = {"device": device, "dtype": PRECISION_TO_TYPE[args.precision]}
|
||||
in_channels = args.latent_channels
|
||||
out_channels = args.latent_channels
|
||||
|
||||
@@ -100,6 +88,8 @@ 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 ========================
|
||||
@@ -114,23 +104,19 @@ 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)
|
||||
crop_start = PROMPT_TEMPLATE[args.prompt_template].get("crop_start", 0)
|
||||
else:
|
||||
crop_start = 0
|
||||
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)
|
||||
if args.prompt_template_video is not None else None)
|
||||
|
||||
text_encoder = TextEncoder(
|
||||
text_encoder_type=args.text_encoder,
|
||||
@@ -184,23 +170,17 @@ class Inference(object):
|
||||
model_path = dit_weight / f"pytorch_model_{load_key}.pt"
|
||||
bare_model = True
|
||||
elif any(str(f).endswith("_model_states.pt") for f in files):
|
||||
files = [
|
||||
f for f in files if str(f).endswith("_model_states.pt")
|
||||
]
|
||||
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"))
|
||||
@@ -210,23 +190,17 @@ class Inference(object):
|
||||
model_path = dit_weight / f"pytorch_model_{load_key}.pt"
|
||||
bare_model = True
|
||||
elif any(str(f).endswith("_model_states.pt") for f in files):
|
||||
files = [
|
||||
f for f in files if str(f).endswith("_model_states.pt")
|
||||
]
|
||||
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"
|
||||
@@ -241,21 +215,18 @@ 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}")
|
||||
|
||||
if bare_model == "unknown" and ("ema" in state_dict
|
||||
or "module" in state_dict):
|
||||
if bare_model == "unknown" and ("ema" in state_dict or "module" in state_dict):
|
||||
bare_model = False
|
||||
if bare_model is False:
|
||||
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
|
||||
|
||||
@@ -264,13 +235,11 @@ class Inference(object):
|
||||
if isinstance(size, int):
|
||||
size = [size]
|
||||
if not isinstance(size, (list, tuple)):
|
||||
raise ValueError(
|
||||
f"Size must be an integer or (height, width), got {size}.")
|
||||
raise ValueError(f"Size must be an integer or (height, width), got {size}.")
|
||||
if len(size) == 1:
|
||||
size = [size[0], size[0]]
|
||||
if len(size) != 2:
|
||||
raise ValueError(
|
||||
f"Size must be an integer or (height, width), got {size}.")
|
||||
raise ValueError(f"Size must be an integer or (height, width), got {size}.")
|
||||
return size
|
||||
|
||||
|
||||
@@ -369,7 +338,7 @@ class HunyuanVideoSampler(Inference):
|
||||
embedded_guidance_scale=None,
|
||||
batch_size=1,
|
||||
num_videos_per_prompt=1,
|
||||
only_save_latent=False,
|
||||
mask_strategy=None,
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
@@ -396,36 +365,22 @@ 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
|
||||
]
|
||||
generator = [torch.Generator("cpu").manual_seed(seed) for seed in seeds]
|
||||
out_dict["seeds"] = seeds
|
||||
|
||||
# ========================================================================
|
||||
@@ -436,13 +391,9 @@ 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)
|
||||
@@ -454,17 +405,14 @@ class HunyuanVideoSampler(Inference):
|
||||
# Arguments: prompt, new_prompt, negative_prompt
|
||||
# ========================================================================
|
||||
if not isinstance(prompt, str):
|
||||
raise TypeError(
|
||||
f"`prompt` must be a string, but got {type(prompt)}")
|
||||
raise TypeError(f"`prompt` must be a string, but got {type(prompt)}")
|
||||
prompt = [prompt.strip()]
|
||||
|
||||
# negative prompt
|
||||
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()]
|
||||
|
||||
# ========================================================================
|
||||
@@ -478,11 +426,9 @@ class HunyuanVideoSampler(Inference):
|
||||
self.pipeline.scheduler = scheduler
|
||||
|
||||
if "884" in self.args.vae:
|
||||
latents_size = [(video_length - 1) // 4 + 1, height // 8,
|
||||
width // 8]
|
||||
latents_size = [(video_length - 1) // 4 + 1, height // 8, width // 8]
|
||||
elif "888" in self.args.vae:
|
||||
latents_size = [(video_length - 1) // 8 + 1, height // 8,
|
||||
width // 8]
|
||||
latents_size = [(video_length - 1) // 8 + 1, height // 8, width // 8]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
|
||||
# ========================================================================
|
||||
@@ -525,10 +471,8 @@ class HunyuanVideoSampler(Inference):
|
||||
vae_ver=self.args.vae,
|
||||
enable_tiling=self.args.vae_tiling,
|
||||
enable_vae_sp=self.args.vae_sp,
|
||||
only_save_latent=only_save_latent,
|
||||
mask_strategy=mask_strategy,
|
||||
)[0]
|
||||
if only_save_latent:
|
||||
return samples
|
||||
out_dict["samples"] = samples
|
||||
out_dict["prompts"] = prompt
|
||||
|
||||
|
||||
@@ -1,10 +1,16 @@
|
||||
import torch
|
||||
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.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
|
||||
|
||||
def attention(
|
||||
@@ -21,23 +27,43 @@ 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 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)
|
||||
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):
|
||||
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)
|
||||
@@ -46,8 +72,7 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask):
|
||||
|
||||
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)
|
||||
@@ -57,28 +82,37 @@ def parallel_attention(q, k, v, img_q_len, img_kv_len, text_mask):
|
||||
sequence_length = query.size(1)
|
||||
encoder_sequence_length = encoder_query.size(1)
|
||||
|
||||
# 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)
|
||||
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)
|
||||
|
||||
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)
|
||||
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)
|
||||
|
||||
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 = 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)
|
||||
|
||||
|
||||
@@ -45,8 +45,7 @@ class PatchEmbed(nn.Module):
|
||||
bias=bias,
|
||||
**factory_kwargs,
|
||||
)
|
||||
nn.init.xavier_uniform_(
|
||||
self.proj.weight.view(self.proj.weight.size(0), -1))
|
||||
nn.init.xavier_uniform_(self.proj.weight.view(self.proj.weight.size(0), -1))
|
||||
if bias:
|
||||
nn.init.zeros_(self.proj.bias)
|
||||
|
||||
@@ -67,12 +66,7 @@ class TextProjection(nn.Module):
|
||||
Adapted from https://github.com/PixArt-alpha/PixArt-alpha/blob/master/diffusion/model/nets/PixArt_blocks.py
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
in_channels,
|
||||
hidden_size,
|
||||
act_layer,
|
||||
dtype=None,
|
||||
device=None):
|
||||
def __init__(self, in_channels, hidden_size, act_layer, dtype=None, device=None):
|
||||
factory_kwargs = {"dtype": dtype, "device": device}
|
||||
super().__init__()
|
||||
self.linear_1 = nn.Linear(
|
||||
@@ -111,14 +105,12 @@ 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) /
|
||||
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:
|
||||
embedding = torch.cat(
|
||||
[embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
return embedding
|
||||
|
||||
|
||||
@@ -145,10 +137,7 @@ 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),
|
||||
)
|
||||
@@ -156,8 +145,6 @@ 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
|
||||
|
||||
@@ -32,21 +32,13 @@ class MLP(nn.Module):
|
||||
hidden_channels = hidden_channels or in_channels
|
||||
bias = to_2tuple(bias)
|
||||
drop_probs = to_2tuple(drop)
|
||||
linear_layer = partial(nn.Conv2d,
|
||||
kernel_size=1) if use_conv else nn.Linear
|
||||
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):
|
||||
@@ -66,15 +58,9 @@ class MLPEmbedder(nn.Module):
|
||||
def __init__(self, in_dim: int, hidden_dim: int, device=None, dtype=None):
|
||||
factory_kwargs = {"device": device, "dtype": dtype}
|
||||
super().__init__()
|
||||
self.in_layer = nn.Linear(in_dim,
|
||||
hidden_dim,
|
||||
bias=True,
|
||||
**factory_kwargs)
|
||||
self.in_layer = nn.Linear(in_dim, hidden_dim, bias=True, **factory_kwargs)
|
||||
self.silu = nn.SiLU()
|
||||
self.out_layer = nn.Linear(hidden_dim,
|
||||
hidden_dim,
|
||||
bias=True,
|
||||
**factory_kwargs)
|
||||
self.out_layer = nn.Linear(hidden_dim, hidden_dim, bias=True, **factory_kwargs)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.out_layer(self.silu(self.in_layer(x)))
|
||||
@@ -83,21 +69,12 @@ 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,
|
||||
@@ -117,10 +94,7 @@ class FinalLayer(nn.Module):
|
||||
# Here we don't distinguish between the modulate types. Just use the simple one.
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
act_layer(),
|
||||
nn.Linear(hidden_size,
|
||||
2 * hidden_size,
|
||||
bias=True,
|
||||
**factory_kwargs),
|
||||
nn.Linear(hidden_size, 2 * hidden_size, bias=True, **factory_kwargs),
|
||||
)
|
||||
# Zero-initialize the modulation
|
||||
nn.init.zeros_(self.adaLN_modulation[1].weight)
|
||||
|
||||
@@ -6,8 +6,7 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.models import ModelMixin
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.models.hunyuan.modules.posemb_layers import \
|
||||
get_nd_rotary_pos_embed
|
||||
from fastvideo.models.hunyuan.modules.posemb_layers import get_nd_rotary_pos_embed
|
||||
from fastvideo.utils.parallel_states import nccl_info
|
||||
|
||||
from .activation_layers import get_activation_layer
|
||||
@@ -53,31 +52,17 @@ 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)
|
||||
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)
|
||||
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_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,
|
||||
@@ -92,30 +77,16 @@ 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)
|
||||
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)
|
||||
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_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,
|
||||
@@ -138,6 +109,7 @@ 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,
|
||||
@@ -158,14 +130,9 @@ 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)
|
||||
@@ -175,34 +142,23 @@ 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),
|
||||
shrink_head(freqs_cis[1], dim=0),
|
||||
)
|
||||
|
||||
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}"
|
||||
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}"
|
||||
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)
|
||||
@@ -214,6 +170,7 @@ 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
|
||||
@@ -221,27 +178,18 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
img_attn, txt_attn = attn[:, :img.shape[1]], attn[:, img.shape[1]:]
|
||||
|
||||
# Calculate the img blocks.
|
||||
img = img + apply_gate(self.img_attn_proj(img_attn),
|
||||
gate=img_mod1_gate)
|
||||
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.
|
||||
txt = txt + apply_gate(self.txt_attn_proj(txt_attn),
|
||||
gate=txt_mod1_gate)
|
||||
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
|
||||
|
||||
|
||||
@@ -277,24 +225,17 @@ class MMSingleStreamBlock(nn.Module):
|
||||
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)
|
||||
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)
|
||||
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(
|
||||
@@ -318,17 +259,13 @@ 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)
|
||||
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
|
||||
|
||||
# Apply QK-Norm if needed.
|
||||
q = self.q_norm(q).to(v)
|
||||
@@ -336,22 +273,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}"
|
||||
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}"
|
||||
img_q, img_k = img_qq, img_kk
|
||||
|
||||
attn = parallel_attention(
|
||||
@@ -361,6 +295,7 @@ 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
|
||||
@@ -463,19 +398,15 @@ 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":
|
||||
@@ -494,21 +425,16 @@ 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)
|
||||
self.guidance_in = (TimestepEmbedder(self.hidden_size, get_activation_layer("silu"), **factory_kwargs)
|
||||
if guidance_embed else None)
|
||||
|
||||
# double blocks
|
||||
@@ -564,12 +490,8 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
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"
|
||||
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"
|
||||
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
|
||||
rope_dim_list,
|
||||
rope_sizes,
|
||||
@@ -592,6 +514,7 @@ 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,
|
||||
@@ -599,15 +522,14 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
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)
|
||||
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))]
|
||||
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]
|
||||
text_states_2 = encoder_hidden_states[:, 0, :self.config.text_states_dim_2]
|
||||
_, _, ot, oh, ow = x.shape # codespell:ignore
|
||||
tt, th, tw = (
|
||||
ot // self.patch_size[0], # codespell:ignore
|
||||
@@ -625,9 +547,7 @@ 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)
|
||||
@@ -637,36 +557,33 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
if self.text_projection == "linear":
|
||||
txt = self.txt_in(txt)
|
||||
elif self.text_projection == "single_refiner":
|
||||
txt = self.txt_in(txt, t,
|
||||
text_mask if self.use_attention_mask else None)
|
||||
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 _, block in enumerate(self.single_blocks):
|
||||
for index, 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, ...])
|
||||
@@ -674,8 +591,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
img = x[:, :img_seq_len, ...]
|
||||
|
||||
# ---------------------------- Final layer ------------------------------
|
||||
img = self.final_layer(img,
|
||||
vec) # (N, T, patch_size ** 2 * out_channels)
|
||||
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."
|
||||
@@ -704,18 +620,18 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
|
||||
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())
|
||||
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())
|
||||
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":
|
||||
|
||||
@@ -18,10 +18,7 @@ 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)
|
||||
@@ -152,5 +149,4 @@ def get_norm_layer(norm_layer):
|
||||
elif norm_layer == "rms":
|
||||
return RMSNorm
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f"Norm layer {norm_layer} is not implemented")
|
||||
raise NotImplementedError(f"Norm layer {norm_layer} is not implemented")
|
||||
|
||||
@@ -75,5 +75,4 @@ def get_norm_layer(norm_layer):
|
||||
elif norm_layer == "rms":
|
||||
return RMSNorm
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
f"Norm layer {norm_layer} is not implemented")
|
||||
raise NotImplementedError(f"Norm layer {norm_layer} is not implemented")
|
||||
|
||||
@@ -100,19 +100,13 @@ 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],
|
||||
x.shape[-1],
|
||||
), f"freqs_cis shape {freqs_cis[0].shape} does not match x shape {x.shape}"
|
||||
shape = [
|
||||
d if i == 1 or i == ndim - 1 else 1
|
||||
for i, d in enumerate(x.shape)
|
||||
]
|
||||
shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
|
||||
return freqs_cis[0].view(*shape), freqs_cis[1].view(*shape)
|
||||
else:
|
||||
# freqs_cis: values in complex space
|
||||
@@ -121,25 +115,18 @@ 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],
|
||||
x.shape[-1],
|
||||
), f"freqs_cis shape {freqs_cis.shape} does not match x shape {x.shape}"
|
||||
shape = [
|
||||
d if i == 1 or i == ndim - 1 else 1
|
||||
for i, d in enumerate(x.shape)
|
||||
]
|
||||
shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
|
||||
return freqs_cis.view(*shape)
|
||||
|
||||
|
||||
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)
|
||||
|
||||
|
||||
@@ -177,15 +164,12 @@ 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
|
||||
@@ -219,28 +203,21 @@ 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):
|
||||
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:
|
||||
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):
|
||||
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:
|
||||
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 = []
|
||||
@@ -300,8 +277,7 @@ def get_1d_rotary_pos_embed(
|
||||
if theta_rescale_factor != 1.0:
|
||||
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:
|
||||
@@ -309,6 +285,5 @@ 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
|
||||
|
||||
@@ -33,30 +33,16 @@ 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)
|
||||
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)
|
||||
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_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,
|
||||
@@ -68,10 +54,7 @@ class IndividualTokenRefinerBlock(nn.Module):
|
||||
|
||||
self.adaLN_modulation = nn.Sequential(
|
||||
act_layer(),
|
||||
nn.Linear(hidden_size,
|
||||
2 * hidden_size,
|
||||
bias=True,
|
||||
**factory_kwargs),
|
||||
nn.Linear(hidden_size, 2 * hidden_size, bias=True, **factory_kwargs),
|
||||
)
|
||||
# Zero-initialize the modulation
|
||||
nn.init.zeros_(self.adaLN_modulation[1].weight)
|
||||
@@ -80,18 +63,14 @@ class IndividualTokenRefinerBlock(nn.Module):
|
||||
def forward(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
c: torch.
|
||||
Tensor, # timestep_aware_representations + context_aware_representations
|
||||
c: torch.Tensor, # timestep_aware_representations + context_aware_representations
|
||||
attn_mask: torch.Tensor = None,
|
||||
):
|
||||
gate_msa, gate_mlp = self.adaLN_modulation(c).chunk(2, dim=1)
|
||||
|
||||
norm_x = self.norm1(x)
|
||||
qkv = self.self_attn_qkv(norm_x)
|
||||
q, k, v = rearrange(qkv,
|
||||
"B L (K H D) -> K B L H D",
|
||||
K=3,
|
||||
H=self.heads_num)
|
||||
q, k, v = rearrange(qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
|
||||
# Apply QK-Norm if needed
|
||||
q = self.self_attn_q_norm(q).to(v)
|
||||
k = self.self_attn_k_norm(k).to(v)
|
||||
@@ -179,18 +158,13 @@ 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)
|
||||
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,
|
||||
@@ -217,10 +191,8 @@ 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 = self.c_embedder(
|
||||
context_aware_representations)
|
||||
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
|
||||
|
||||
x = self.input_embedder(x)
|
||||
|
||||
@@ -23,24 +23,20 @@ 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}")
|
||||
# from_pretrained will ensure that the model is in eval mode.
|
||||
|
||||
if text_encoder_precision is not None:
|
||||
text_encoder = text_encoder.to(
|
||||
dtype=PRECISION_TO_TYPE[text_encoder_precision])
|
||||
text_encoder = text_encoder.to(dtype=PRECISION_TO_TYPE[text_encoder_precision])
|
||||
|
||||
text_encoder.requires_grad_(False)
|
||||
|
||||
@@ -53,22 +49,16 @@ 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:
|
||||
logger.info(
|
||||
f"Loading tokenizer ({tokenizer_type}) from: {tokenizer_path}")
|
||||
logger.info(f"Loading tokenizer ({tokenizer_type}) from: {tokenizer_path}")
|
||||
|
||||
if tokenizer_type == "clipL":
|
||||
tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path,
|
||||
max_length=77)
|
||||
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}")
|
||||
|
||||
@@ -125,16 +115,12 @@ 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
|
||||
@@ -144,10 +130,8 @@ 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']}")
|
||||
@@ -156,8 +140,7 @@ class TextEncoder(nn.Module):
|
||||
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, "
|
||||
@@ -170,8 +153,7 @@ class TextEncoder(nn.Module):
|
||||
elif "llm" in text_encoder_type or "glm" in text_encoder_type:
|
||||
self.output_key = output_key or "last_hidden_state"
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported text encoder type: {text_encoder_type}")
|
||||
raise ValueError(f"Unsupported text encoder type: {text_encoder_type}")
|
||||
|
||||
self.model, self.model_path = load_text_encoder(
|
||||
text_encoder_type=self.text_encoder_type,
|
||||
@@ -226,10 +208,7 @@ 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):
|
||||
@@ -262,8 +241,7 @@ class TextEncoder(nn.Module):
|
||||
**kwargs,
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported tokenize_input_type: {tokenize_input_type}")
|
||||
raise ValueError(f"Unsupported tokenize_input_type: {tokenize_input_type}")
|
||||
|
||||
def encode(
|
||||
self,
|
||||
@@ -291,27 +269,21 @@ class TextEncoder(nn.Module):
|
||||
return_texts (bool): Whether to return the decoded texts. Defaults to False.
|
||||
"""
|
||||
device = self.model.device if device is None else device
|
||||
use_attention_mask = use_default(use_attention_mask,
|
||||
self.use_attention_mask)
|
||||
hidden_state_skip_layer = use_default(hidden_state_skip_layer,
|
||||
self.hidden_state_skip_layer)
|
||||
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)
|
||||
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)]
|
||||
last_hidden_state = outputs.hidden_states[-(hidden_state_skip_layer + 1)]
|
||||
# Real last hidden state already has layer norm applied. So here we only apply it
|
||||
# for intermediate layers.
|
||||
if hidden_state_skip_layer > 0 and self.apply_final_norm:
|
||||
last_hidden_state = self.model.final_layer_norm(
|
||||
last_hidden_state)
|
||||
last_hidden_state = self.model.final_layer_norm(last_hidden_state)
|
||||
else:
|
||||
last_hidden_state = outputs[self.output_key]
|
||||
|
||||
@@ -325,12 +297,10 @@ 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(
|
||||
|
||||
@@ -45,11 +45,7 @@ def safe_file(path):
|
||||
return path
|
||||
|
||||
|
||||
def save_videos_grid(videos: torch.Tensor,
|
||||
path: str,
|
||||
rescale=False,
|
||||
n_rows=1,
|
||||
fps=24):
|
||||
def save_videos_grid(videos: torch.Tensor, path: str, rescale=False, n_rows=1, fps=24):
|
||||
"""save videos by video tensor
|
||||
copy from https://github.com/guoyww/AnimateDiff/blob/e92bd5671ba62c0d774a32951453e328018b7c5b/animatediff/utils/util.py#L61
|
||||
|
||||
|
||||
@@ -31,8 +31,7 @@ def load_vae(
|
||||
logger.info(f"Loading 3D VAE model ({vae_type}) from: {vae_path}")
|
||||
config = AutoencoderKLCausal3D.load_config(vae_path)
|
||||
if sample_size:
|
||||
vae = AutoencoderKLCausal3D.from_config(config,
|
||||
sample_size=sample_size)
|
||||
vae = AutoencoderKLCausal3D.from_config(config, sample_size=sample_size)
|
||||
else:
|
||||
vae = AutoencoderKLCausal3D.from_config(config)
|
||||
|
||||
@@ -43,10 +42,7 @@ 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
|
||||
|
||||
@@ -35,15 +35,13 @@ except ImportError:
|
||||
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)
|
||||
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 BaseOutput, DecoderCausal3D, DecoderOutput, DiagonalGaussianDistribution, EncoderCausal3D
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -113,12 +111,8 @@ 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
|
||||
@@ -130,11 +124,9 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
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):
|
||||
@@ -207,12 +199,10 @@ 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)
|
||||
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
|
||||
|
||||
return processors
|
||||
|
||||
@@ -244,21 +234,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):
|
||||
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)
|
||||
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_attn_processor(name, module, processor)
|
||||
@@ -268,11 +254,9 @@ 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(
|
||||
@@ -282,11 +266,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.
|
||||
|
||||
@@ -304,9 +286,8 @@ 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:
|
||||
@@ -323,11 +304,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
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:
|
||||
@@ -336,9 +313,8 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
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)
|
||||
@@ -369,9 +345,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
"""
|
||||
if self.use_slicing and z.shape[0] > 1:
|
||||
decoded_slices = [
|
||||
self._decode(z_slice).sample for z_slice in z.split(1)
|
||||
]
|
||||
decoded_slices = [self._decode(z_slice).sample for z_slice in z.split(1)]
|
||||
decoded = torch.cat(decoded_slices)
|
||||
else:
|
||||
decoded = self._decode(z).sample
|
||||
@@ -381,28 +355,26 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
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(
|
||||
@@ -429,10 +401,8 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
If return_dict is True, a [`~models.autoencoder_kl.AutoencoderKLOutput`] is returned, otherwise a plain
|
||||
`tuple` is returned.
|
||||
"""
|
||||
overlap_size = int(self.tile_sample_min_size *
|
||||
(1 - self.tile_overlap_factor))
|
||||
blend_extent = int(self.tile_latent_min_size *
|
||||
self.tile_overlap_factor)
|
||||
overlap_size = int(self.tile_sample_min_size * (1 - self.tile_overlap_factor))
|
||||
blend_extent = int(self.tile_latent_min_size * self.tile_overlap_factor)
|
||||
row_limit = self.tile_latent_min_size - blend_extent
|
||||
|
||||
# Split video into tiles and encode them separately.
|
||||
@@ -440,8 +410,7 @@ 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)
|
||||
@@ -471,8 +440,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
def spatial_tiled_decode(self,
|
||||
z: torch.FloatTensor,
|
||||
return_dict: bool = True
|
||||
) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
r"""
|
||||
Decode a batch of images/videos using a tiled decoder.
|
||||
|
||||
@@ -486,10 +454,8 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
If return_dict is True, a [`~models.vae.DecoderOutput`] is returned, otherwise a plain `tuple` is
|
||||
returned.
|
||||
"""
|
||||
overlap_size = int(self.tile_latent_min_size *
|
||||
(1 - self.tile_overlap_factor))
|
||||
blend_extent = int(self.tile_sample_min_size *
|
||||
self.tile_overlap_factor)
|
||||
overlap_size = int(self.tile_latent_min_size * (1 - self.tile_overlap_factor))
|
||||
blend_extent = int(self.tile_sample_min_size * self.tile_overlap_factor)
|
||||
row_limit = self.tile_sample_min_size - blend_extent
|
||||
|
||||
# Split z into overlapping tiles and decode them separately.
|
||||
@@ -498,8 +464,7 @@ 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)
|
||||
@@ -523,24 +488,19 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
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))
|
||||
blend_extent = int(self.tile_latent_min_tsize *
|
||||
self.tile_overlap_factor)
|
||||
overlap_size = int(self.tile_sample_min_tsize * (1 - self.tile_overlap_factor))
|
||||
blend_extent = int(self.tile_latent_min_tsize * self.tile_overlap_factor)
|
||||
t_limit = self.tile_latent_min_tsize - blend_extent
|
||||
|
||||
# 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):
|
||||
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)
|
||||
@@ -566,25 +526,20 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
def temporal_tiled_decode(self,
|
||||
z: torch.FloatTensor,
|
||||
return_dict: bool = True
|
||||
) -> Union[DecoderOutput, 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
|
||||
overlap_size = int(self.tile_latent_min_tsize *
|
||||
(1 - self.tile_overlap_factor))
|
||||
blend_extent = int(self.tile_sample_min_tsize *
|
||||
self.tile_overlap_factor)
|
||||
overlap_size = int(self.tile_latent_min_tsize * (1 - self.tile_overlap_factor))
|
||||
blend_extent = int(self.tile_sample_min_tsize * self.tile_overlap_factor)
|
||||
t_limit = self.tile_sample_min_tsize - blend_extent
|
||||
|
||||
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):
|
||||
decoded = self.spatial_tiled_decode(tile,
|
||||
return_dict=True).sample
|
||||
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)
|
||||
decoded = self.decoder(tile)
|
||||
@@ -605,22 +560,19 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
return DecoderOutput(sample=dec)
|
||||
|
||||
def _parallel_data_generator(self, gathered_results,
|
||||
gathered_dim_metadata):
|
||||
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)
|
||||
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]:
|
||||
return_dict: bool = True) -> Union[DecoderOutput, torch.FloatTensor]:
|
||||
"""
|
||||
Parallel version of tiled_decode that distributes both temporal and spatial computation across GPUs
|
||||
"""
|
||||
@@ -628,16 +580,12 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
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_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_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
|
||||
@@ -655,8 +603,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
local_results = []
|
||||
local_dim_metadata = []
|
||||
# Process assigned tiles
|
||||
for local_idx, global_idx in enumerate(
|
||||
range(start_tile_idx, end_tile_idx)):
|
||||
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
|
||||
@@ -670,8 +617,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
|
||||
# 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]
|
||||
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)
|
||||
@@ -691,13 +637,8 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
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)
|
||||
]
|
||||
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)
|
||||
@@ -707,16 +648,13 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
# 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
|
||||
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):
|
||||
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
|
||||
@@ -726,11 +664,9 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
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)
|
||||
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)
|
||||
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, :, :])
|
||||
@@ -748,8 +684,7 @@ class AutoencoderKLCausal3D(ModelMixin, ConfigMixin, FromOriginalVAEMixin):
|
||||
result_row = []
|
||||
for j, tile in enumerate(row):
|
||||
if i > 0:
|
||||
tile = self.blend_v(spatial_rows[i - 1][j], tile,
|
||||
blend_extent)
|
||||
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])
|
||||
@@ -806,9 +741,7 @@ 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
|
||||
|
||||
|
||||
@@ -31,16 +31,9 @@ from torch import nn
|
||||
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)
|
||||
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
|
||||
@@ -78,12 +71,7 @@ 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)
|
||||
@@ -135,10 +123,7 @@ 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
|
||||
@@ -175,14 +160,10 @@ 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
|
||||
@@ -260,15 +241,11 @@ 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
|
||||
|
||||
@@ -325,58 +302,36 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
groups_out = groups
|
||||
|
||||
if self.time_embedding_norm == "ada_group":
|
||||
self.norm1 = AdaGroupNorm(temb_channels,
|
||||
in_channels,
|
||||
groups,
|
||||
eps=eps)
|
||||
self.norm1 = AdaGroupNorm(temb_channels, in_channels, groups, eps=eps)
|
||||
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)
|
||||
self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, stride=1)
|
||||
|
||||
if temb_channels is not None:
|
||||
if self.time_embedding_norm == "default":
|
||||
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"):
|
||||
self.time_emb_proj = linear_cls(temb_channels, 2 * out_channels)
|
||||
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
|
||||
|
||||
if self.time_embedding_norm == "ada_group":
|
||||
self.norm2 = AdaGroupNorm(temb_channels,
|
||||
out_channels,
|
||||
groups_out,
|
||||
eps=eps)
|
||||
self.norm2 = AdaGroupNorm(temb_channels, out_channels, groups_out, eps=eps)
|
||||
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)
|
||||
|
||||
@@ -384,12 +339,10 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
if self.up:
|
||||
self.upsample = UpsampleCausal3D(in_channels, use_conv=False)
|
||||
elif self.down:
|
||||
self.downsample = DownsampleCausal3D(in_channels,
|
||||
use_conv=False,
|
||||
name="op")
|
||||
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:
|
||||
@@ -409,8 +362,7 @@ 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)
|
||||
@@ -438,8 +390,7 @@ 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)
|
||||
@@ -456,8 +407,7 @@ class ResnetBlockCausal3D(nn.Module):
|
||||
if self.conv_shortcut is not None:
|
||||
input_tensor = self.conv_shortcut(input_tensor)
|
||||
|
||||
output_tensor = (input_tensor +
|
||||
hidden_states) / self.output_scale_factor
|
||||
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
|
||||
|
||||
return output_tensor
|
||||
|
||||
@@ -497,9 +447,7 @@ 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,
|
||||
@@ -553,8 +501,7 @@ 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,
|
||||
@@ -595,13 +542,11 @@ 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 = [
|
||||
@@ -636,9 +581,7 @@ 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,
|
||||
@@ -664,29 +607,19 @@ class UNetMidBlockCausal3D(nn.Module):
|
||||
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)
|
||||
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)
|
||||
hidden_states = resnet(hidden_states, temb)
|
||||
|
||||
return hidden_states
|
||||
@@ -745,9 +678,7 @@ class DownEncoderBlockCausal3D(nn.Module):
|
||||
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)
|
||||
|
||||
@@ -761,21 +692,21 @@ 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 = []
|
||||
|
||||
@@ -8,8 +8,7 @@ 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 .unet_causal_3d_blocks import CausalConv3d, UNetMidBlockCausal3D, get_down_block3d, get_up_block3d
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -47,10 +46,7 @@ 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([])
|
||||
|
||||
@@ -60,33 +56,25 @@ class EncoderCausal3D(nn.Module):
|
||||
input_channel = output_channel
|
||||
output_channel = block_out_channels[i]
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
num_spatial_downsample_layers = int(
|
||||
np.log2(spatial_compression_ratio))
|
||||
num_spatial_downsample_layers = int(np.log2(spatial_compression_ratio))
|
||||
num_time_downsample_layers = int(np.log2(time_compression_ratio))
|
||||
|
||||
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_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)
|
||||
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 = tuple(downsample_stride_T +
|
||||
downsample_stride_HW)
|
||||
downsample_stride = tuple(downsample_stride_T + downsample_stride_HW)
|
||||
down_block = get_down_block3d(
|
||||
down_block_type,
|
||||
num_layers=self.layers_per_block,
|
||||
in_channels=input_channel,
|
||||
out_channels=output_channel,
|
||||
add_downsample=bool(add_spatial_downsample
|
||||
or add_time_downsample),
|
||||
add_downsample=bool(add_spatial_downsample or add_time_downsample),
|
||||
downsample_stride=downsample_stride,
|
||||
resnet_eps=1e-6,
|
||||
downsample_padding=0,
|
||||
@@ -111,20 +99,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."""
|
||||
assert len(
|
||||
sample.shape) == 5, "The input tensor should have 5 dimensions"
|
||||
assert len(sample.shape) == 5, "The input tensor should have 5 dimensions"
|
||||
|
||||
sample = self.conv_in(sample)
|
||||
|
||||
@@ -165,10 +148,7 @@ 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([])
|
||||
|
||||
@@ -180,8 +160,7 @@ class DecoderCausal3D(nn.Module):
|
||||
resnet_eps=1e-6,
|
||||
resnet_act_fn=act_fn,
|
||||
output_scale_factor=1,
|
||||
resnet_time_scale_shift="default"
|
||||
if norm_type == "group" else norm_type,
|
||||
resnet_time_scale_shift="default" if norm_type == "group" else norm_type,
|
||||
attention_head_dim=block_out_channels[-1],
|
||||
resnet_groups=norm_num_groups,
|
||||
temb_channels=temb_channels,
|
||||
@@ -195,25 +174,19 @@ class DecoderCausal3D(nn.Module):
|
||||
prev_output_channel = output_channel
|
||||
output_channel = reversed_block_out_channels[i]
|
||||
is_final_block = i == len(block_out_channels) - 1
|
||||
num_spatial_upsample_layers = int(
|
||||
np.log2(spatial_compression_ratio))
|
||||
num_spatial_upsample_layers = int(np.log2(spatial_compression_ratio))
|
||||
num_time_upsample_layers = int(np.log2(time_compression_ratio))
|
||||
|
||||
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_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 = tuple(upsample_scale_factor_T + upsample_scale_factor_HW)
|
||||
up_block = get_up_block3d(
|
||||
up_block_type,
|
||||
num_layers=self.layers_per_block + 1,
|
||||
@@ -234,17 +207,11 @@ class DecoderCausal3D(nn.Module):
|
||||
|
||||
# out
|
||||
if norm_type == "spatial":
|
||||
self.conv_norm_out = SpatialNorm(block_out_channels[0],
|
||||
temb_channels)
|
||||
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)
|
||||
self.conv_out = CausalConv3d(block_out_channels[0], out_channels, kernel_size=3)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
@@ -254,8 +221,7 @@ class DecoderCausal3D(nn.Module):
|
||||
latent_embeds: Optional[torch.FloatTensor] = None,
|
||||
) -> torch.FloatTensor:
|
||||
r"""The forward method of the `DecoderCausal3D` class."""
|
||||
assert len(
|
||||
sample.shape) == 5, "The input tensor should have 5 dimensions."
|
||||
assert len(sample.shape) == 5, "The input tensor should have 5 dimensions."
|
||||
|
||||
sample = self.conv_in(sample)
|
||||
|
||||
@@ -289,15 +255,12 @@ 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)
|
||||
@@ -334,14 +297,11 @@ 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:
|
||||
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
|
||||
sample = randn_tensor(
|
||||
self.mean.shape,
|
||||
@@ -364,20 +324,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)
|
||||
return 0.5 * torch.sum(
|
||||
logtwopi + self.logvar +
|
||||
torch.pow(sample - self.mean, 2) / self.var,
|
||||
logtwopi + self.logvar + torch.pow(sample - self.mean, 2) / self.var,
|
||||
dim=dims,
|
||||
)
|
||||
|
||||
|
||||
@@ -21,29 +21,23 @@ 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.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 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)
|
||||
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)
|
||||
return encoder_state.narrow(dim, nccl_info.rank_within_group * local_heads, local_heads)
|
||||
|
||||
|
||||
class HunyuanVideoAttnProcessor2_0:
|
||||
@@ -51,8 +45,7 @@ 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."
|
||||
)
|
||||
"HunyuanVideoAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
@@ -66,8 +59,7 @@ class HunyuanVideoAttnProcessor2_0:
|
||||
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)
|
||||
hidden_states = torch.cat([hidden_states, encoder_hidden_states], dim=1)
|
||||
|
||||
# 1. QKV projections
|
||||
query = attn.to_q(hidden_states)
|
||||
@@ -96,18 +88,14 @@ class HunyuanVideoAttnProcessor2_0:
|
||||
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),
|
||||
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),
|
||||
apply_rotary_emb(key[:, :, :-encoder_hidden_states.shape[1]], image_rotary_emb),
|
||||
key[:, :, -encoder_hidden_states.shape[1]:],
|
||||
],
|
||||
dim=2,
|
||||
@@ -122,16 +110,12 @@ class HunyuanVideoAttnProcessor2_0:
|
||||
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)
|
||||
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)
|
||||
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)
|
||||
|
||||
@@ -140,17 +124,10 @@ class HunyuanVideoAttnProcessor2_0:
|
||||
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) #
|
||||
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)
|
||||
|
||||
@@ -171,24 +148,15 @@ class HunyuanVideoAttnProcessor2_0:
|
||||
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)
|
||||
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)
|
||||
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()
|
||||
(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)
|
||||
@@ -225,35 +193,26 @@ class HunyuanVideoPatchEmbed(nn.Module):
|
||||
) -> 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)
|
||||
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
|
||||
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:
|
||||
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]:
|
||||
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)
|
||||
@@ -274,9 +233,7 @@ class HunyuanVideoIndividualTokenRefinerBlock(nn.Module):
|
||||
|
||||
hidden_size = num_attention_heads * attention_head_dim
|
||||
|
||||
self.norm1 = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=True,
|
||||
eps=1e-6)
|
||||
self.norm1 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6)
|
||||
self.attn = Attention(
|
||||
query_dim=hidden_size,
|
||||
cross_attention_dim=None,
|
||||
@@ -285,13 +242,8 @@ class HunyuanVideoIndividualTokenRefinerBlock(nn.Module):
|
||||
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.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)
|
||||
|
||||
@@ -352,9 +304,7 @@ class HunyuanVideoIndividualTokenRefiner(nn.Module):
|
||||
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_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
|
||||
@@ -381,8 +331,8 @@ class HunyuanVideoTokenRefiner(nn.Module):
|
||||
|
||||
hidden_size = num_attention_heads * attention_head_dim
|
||||
|
||||
self.time_text_embed = CombinedTimestepTextProjEmbeddings(
|
||||
embedding_dim=hidden_size, pooled_projection_dim=in_channels)
|
||||
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,
|
||||
@@ -404,8 +354,7 @@ class HunyuanVideoTokenRefiner(nn.Module):
|
||||
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 = (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)
|
||||
@@ -417,11 +366,7 @@ class HunyuanVideoTokenRefiner(nn.Module):
|
||||
|
||||
class HunyuanVideoRotaryPosEmbed(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
patch_size: int,
|
||||
patch_size_t: int,
|
||||
rope_dim: List[int],
|
||||
theta: float = 256.0) -> None:
|
||||
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
|
||||
@@ -432,8 +377,7 @@ class HunyuanVideoRotaryPosEmbed(nn.Module):
|
||||
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
|
||||
num_frames * nccl_info.sp_size // self.patch_size_t, height // self.patch_size, width // self.patch_size
|
||||
]
|
||||
|
||||
axes_grids = []
|
||||
@@ -441,26 +385,18 @@ class HunyuanVideoRotaryPosEmbed(nn.Module):
|
||||
# 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)
|
||||
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)
|
||||
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)
|
||||
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
|
||||
|
||||
|
||||
@@ -505,8 +441,7 @@ class HunyuanVideoSingleTransformerBlock(nn.Module):
|
||||
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)
|
||||
hidden_states = torch.cat([hidden_states, encoder_hidden_states], dim=1)
|
||||
|
||||
residual = hidden_states
|
||||
|
||||
@@ -554,8 +489,7 @@ class HunyuanVideoTransformerBlock(nn.Module):
|
||||
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.norm1_context = AdaLayerNormZero(hidden_size, norm_type="layer_norm")
|
||||
|
||||
self.attn = Attention(
|
||||
query_dim=hidden_size,
|
||||
@@ -571,19 +505,11 @@ class HunyuanVideoTransformerBlock(nn.Module):
|
||||
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 = 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")
|
||||
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,
|
||||
@@ -594,8 +520,7 @@ class HunyuanVideoTransformerBlock(nn.Module):
|
||||
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_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)
|
||||
|
||||
@@ -609,30 +534,25 @@ class HunyuanVideoTransformerBlock(nn.Module):
|
||||
|
||||
# 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)
|
||||
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]
|
||||
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
|
||||
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):
|
||||
class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin):
|
||||
r"""
|
||||
A Transformer model for video-like data used in [HunyuanVideo](https://huggingface.co/tencent/HunyuanVideo).
|
||||
|
||||
@@ -699,26 +619,19 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
|
||||
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)
|
||||
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)
|
||||
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)
|
||||
HunyuanVideoTransformerBlock(num_attention_heads, attention_head_dim, mlp_ratio=mlp_ratio, qk_norm=qk_norm)
|
||||
for _ in range(num_layers)
|
||||
])
|
||||
|
||||
@@ -727,17 +640,12 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
|
||||
HunyuanVideoSingleTransformerBlock(num_attention_heads,
|
||||
attention_head_dim,
|
||||
mlp_ratio=mlp_ratio,
|
||||
qk_norm=qk_norm)
|
||||
for _ in range(num_single_layers)
|
||||
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.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
|
||||
|
||||
@@ -752,15 +660,12 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
|
||||
# set recursively
|
||||
processors = {}
|
||||
|
||||
def fn_recursive_add_processors(name: str, module: torch.nn.Module,
|
||||
processors: Dict[str,
|
||||
AttentionProcessor]):
|
||||
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)
|
||||
fn_recursive_add_processors(f"{name}.{sub_name}", child, processors)
|
||||
|
||||
return processors
|
||||
|
||||
@@ -770,9 +675,7 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
|
||||
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]]):
|
||||
def set_attn_processor(self, processor: Union[AttentionProcessor, Dict[str, AttentionProcessor]]):
|
||||
r"""
|
||||
Sets the attention processor to use to compute attention.
|
||||
|
||||
@@ -790,11 +693,9 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
|
||||
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):
|
||||
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)
|
||||
@@ -802,8 +703,7 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
|
||||
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)
|
||||
fn_recursive_attn_processor(f"{name}.{sub_name}", child, processor)
|
||||
|
||||
for name, module in self.named_children():
|
||||
fn_recursive_attn_processor(name, module, processor)
|
||||
@@ -823,9 +723,7 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
|
||||
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)
|
||||
guidance = torch.tensor([6016.0], device=hidden_states.device, dtype=torch.bfloat16)
|
||||
|
||||
if attention_kwargs is not None:
|
||||
attention_kwargs = attention_kwargs.copy()
|
||||
@@ -837,11 +735,8 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
|
||||
# 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."
|
||||
)
|
||||
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
|
||||
@@ -849,8 +744,7 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
|
||||
post_patch_height = height // p
|
||||
post_patch_width = width // p
|
||||
|
||||
pooled_projections = encoder_hidden_states[:, 0, :self.config.
|
||||
pooled_projection_dim]
|
||||
pooled_projections = encoder_hidden_states[:, 0, :self.config.pooled_projection_dim]
|
||||
encoder_hidden_states = encoder_hidden_states[:, 1:]
|
||||
|
||||
# 1. RoPE
|
||||
@@ -859,9 +753,7 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
|
||||
# 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)
|
||||
encoder_hidden_states = self.context_embedder(encoder_hidden_states, timestep, encoder_attention_mask)
|
||||
|
||||
# 3. Attention mask preparation
|
||||
latent_sequence_length = hidden_states.shape[1]
|
||||
@@ -873,13 +765,11 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
|
||||
device=hidden_states.device,
|
||||
dtype=torch.bool) # [B, N, N]
|
||||
|
||||
effective_condition_sequence_length = encoder_attention_mask.sum(
|
||||
dim=1, dtype=torch.int)
|
||||
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
|
||||
attention_mask[i, :effective_sequence_length[i], :effective_sequence_length[i]] = True
|
||||
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
@@ -894,9 +784,7 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
|
||||
|
||||
return custom_forward
|
||||
|
||||
ckpt_kwargs: Dict[str, Any] = {
|
||||
"use_reentrant": False
|
||||
} if is_torch_version(">=", "1.11.0") else {}
|
||||
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(
|
||||
@@ -922,23 +810,19 @@ class HunyuanVideoTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin,
|
||||
|
||||
else:
|
||||
for block in self.transformer_blocks:
|
||||
hidden_states, encoder_hidden_states = block(
|
||||
hidden_states, encoder_hidden_states, temb, attention_mask,
|
||||
image_rotary_emb)
|
||||
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)
|
||||
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.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)
|
||||
|
||||
|
||||
@@ -20,22 +20,18 @@ 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.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 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)
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
@@ -66,14 +62,13 @@ EXAMPLE_DOC_STRING = """
|
||||
"""
|
||||
|
||||
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|>"),
|
||||
"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,
|
||||
}
|
||||
@@ -112,28 +107,22 @@ 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)
|
||||
@@ -195,13 +184,10 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
)
|
||||
|
||||
self.vae_scale_factor_temporal = (self.vae.temporal_compression_ratio
|
||||
if hasattr(self, "vae")
|
||||
and self.vae is not None else 4)
|
||||
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)
|
||||
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,
|
||||
@@ -263,12 +249,9 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
# 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)
|
||||
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
|
||||
|
||||
@@ -295,25 +278,17 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
)
|
||||
|
||||
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}")
|
||||
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
|
||||
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)
|
||||
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, -1)
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
@@ -365,13 +340,10 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
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}."
|
||||
)
|
||||
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):
|
||||
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]}"
|
||||
)
|
||||
@@ -386,28 +358,18 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
" 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)}"
|
||||
)
|
||||
"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)}"
|
||||
)
|
||||
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()}"
|
||||
)
|
||||
f"`prompt_template` has to contain a key `template` but only found {prompt_template.keys()}")
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
@@ -418,8 +380,7 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
num_frames: int = 129,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
generator: Optional[Union[torch.Generator,
|
||||
List[torch.Generator]]] = None,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if latents is not None:
|
||||
@@ -435,13 +396,9 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
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.")
|
||||
|
||||
latents = randn_tensor(shape,
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=dtype)
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
return latents
|
||||
|
||||
def enable_vae_slicing(self):
|
||||
@@ -502,8 +459,7 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
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,
|
||||
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,
|
||||
@@ -511,8 +467,7 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
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,
|
||||
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,
|
||||
@@ -591,8 +546,7 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
indicating whether the corresponding generated image contains "not-safe-for-work" (nsfw) content.
|
||||
"""
|
||||
|
||||
if isinstance(callback_on_step_end,
|
||||
(PipelineCallback, MultiPipelineCallbacks)):
|
||||
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
|
||||
@@ -640,8 +594,7 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
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
|
||||
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,
|
||||
@@ -651,8 +604,7 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = self.transformer.config.in_channels
|
||||
num_latent_frames = (num_frames -
|
||||
1) // self.vae_scale_factor_temporal + 1
|
||||
num_latent_frames = (num_frames - 1) // self.vae_scale_factor_temporal + 1
|
||||
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
@@ -668,19 +620,14 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
# 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 = 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
|
||||
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
|
||||
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:
|
||||
@@ -694,17 +641,14 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
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]),
|
||||
(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)
|
||||
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]
|
||||
encoder_hidden_states=encoder_hidden_states, # [1, 257, 4096]
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
guidance=guidance,
|
||||
@@ -713,37 +657,28 @@ class HunyuanVideoPipeline(DiffusionPipeline, HunyuanVideoLoraLoaderMixin):
|
||||
)[0]
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False)[0]
|
||||
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)
|
||||
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)
|
||||
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):
|
||||
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
|
||||
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)
|
||||
video = self.video_processor.postprocess_video(video, output_type=output_type)
|
||||
else:
|
||||
video = latents
|
||||
|
||||
|
||||
@@ -6,18 +6,9 @@ from safetensors.torch import save_file
|
||||
|
||||
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()
|
||||
|
||||
@@ -39,36 +30,22 @@ 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
|
||||
@@ -77,25 +54,19 @@ 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")
|
||||
@@ -104,18 +75,13 @@ 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")
|
||||
@@ -124,46 +90,34 @@ 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")
|
||||
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.linear.bias"] = original_state_dict.pop(
|
||||
"proj_out.bias")
|
||||
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")
|
||||
new_state_dict["pos_frequencies"] = original_state_dict.pop("pos_frequencies")
|
||||
|
||||
print("Remaining Keys:", original_state_dict.keys())
|
||||
|
||||
@@ -178,271 +132,182 @@ 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")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+1}.stack.0.bias"] = original_state_dict.pop(
|
||||
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")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+1}.stack.2.bias"] = original_state_dict.pop(
|
||||
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")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+1}.stack.3.bias"] = original_state_dict.pop(
|
||||
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")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+1}.stack.5.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_in.resnets.{i}.conv2.conv.bias")
|
||||
encoder_state_dict[f"layers.{i+1}.stack.0.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
encoder_state_dict[f"layers.{i+1}.stack.2.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
encoder_state_dict[f"layers.{i+1}.stack.3.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
encoder_state_dict[f"layers.{i+1}.stack.5.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
|
||||
# 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.bias"] = original_state_dict.pop(
|
||||
f"{prefix}down_blocks.{block}.conv_in.conv.bias")
|
||||
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")
|
||||
|
||||
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")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+7}.stack.0.bias"] = original_state_dict.pop(
|
||||
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")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+7}.stack.2.bias"] = original_state_dict.pop(
|
||||
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")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+7}.stack.3.bias"] = original_state_dict.pop(
|
||||
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")
|
||||
encoder_state_dict[
|
||||
f"layers.{i+7}.stack.5.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.conv2.conv.bias")
|
||||
encoder_state_dict[f"layers.{i+7}.stack.0.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
encoder_state_dict[f"layers.{i+7}.stack.2.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
encoder_state_dict[f"layers.{i+7}.stack.3.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
encoder_state_dict[f"layers.{i+7}.stack.5.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
|
||||
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")
|
||||
v = original_state_dict.pop(
|
||||
f"{prefix}block_out.attentions.{i}.to_v.weight")
|
||||
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")
|
||||
v = original_state_dict.pop(f"{prefix}block_out.attentions.{i}.to_v.weight")
|
||||
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.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")
|
||||
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")
|
||||
decoder_state_dict[
|
||||
f"blocks.0.{i+1}.stack.2.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
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")
|
||||
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")
|
||||
decoder_state_dict[
|
||||
f"blocks.0.{i+1}.stack.5.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
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")
|
||||
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")
|
||||
decoder_state_dict[f"blocks.0.{i+1}.stack.2.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
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")
|
||||
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")
|
||||
decoder_state_dict[f"blocks.0.{i+1}.stack.5.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
|
||||
# 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}.proj.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
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")
|
||||
decoder_state_dict[f"blocks.{block+1}.proj.bias"] = original_state_dict.pop(
|
||||
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")
|
||||
decoder_state_dict[
|
||||
f"blocks.4.{i}.stack.0.bias"] = original_state_dict.pop(
|
||||
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")
|
||||
decoder_state_dict[
|
||||
f"blocks.4.{i}.stack.2.bias"] = original_state_dict.pop(
|
||||
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")
|
||||
decoder_state_dict[
|
||||
f"blocks.4.{i}.stack.3.bias"] = original_state_dict.pop(
|
||||
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")
|
||||
decoder_state_dict[
|
||||
f"blocks.4.{i}.stack.5.bias"] = original_state_dict.pop(
|
||||
f"{prefix}block_out.resnets.{i}.conv2.conv.bias")
|
||||
decoder_state_dict[f"blocks.4.{i}.stack.0.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
decoder_state_dict[f"blocks.4.{i}.stack.2.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
decoder_state_dict[f"blocks.4.{i}.stack.3.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
decoder_state_dict[f"blocks.4.{i}.stack.5.weight"] = original_state_dict.pop(
|
||||
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")
|
||||
|
||||
# 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
|
||||
|
||||
@@ -469,8 +334,7 @@ def main(args):
|
||||
ensure_directory_exists(transformer_path)
|
||||
|
||||
print("Converting transformer model...")
|
||||
transformer_state_dict = convert_diffusers_transformer_to_mochi(
|
||||
pipe.transformer.state_dict())
|
||||
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}")
|
||||
|
||||
@@ -482,8 +346,7 @@ def main(args):
|
||||
ensure_directory_exists(decoder_path)
|
||||
|
||||
print("Converting VAE models...")
|
||||
encoder_state_dict, decoder_state_dict = convert_diffusers_vae_to_mochi(
|
||||
pipe.vae.state_dict())
|
||||
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}")
|
||||
@@ -491,9 +354,7 @@ 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__":
|
||||
|
||||
@@ -21,22 +21,18 @@ from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import PeftAdapterMixin
|
||||
from diffusers.models.attention import FeedForward as HF_FeedForward
|
||||
from diffusers.models.attention_processor import Attention
|
||||
from diffusers.models.embeddings import (MochiCombinedTimestepCaptionEmbedding,
|
||||
PatchEmbed)
|
||||
from diffusers.models.embeddings import MochiCombinedTimestepCaptionEmbedding, PatchEmbed
|
||||
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.utils import USE_PEFT_BACKEND, is_torch_version, logging, scale_lora_layers, unscale_lora_layers
|
||||
from diffusers.utils.torch_utils import maybe_allow_in_graph
|
||||
from liger_kernel.ops.swiglu import LigerSiLUMulFunction
|
||||
|
||||
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
|
||||
from fastvideo.models.mochi_hf.norm import (MochiLayerNormContinuous,
|
||||
MochiModulatedRMSNorm,
|
||||
MochiRMSNorm, MochiRMSNormZero)
|
||||
from fastvideo.models.mochi_hf.norm import (MochiLayerNormContinuous, MochiModulatedRMSNorm, MochiRMSNorm,
|
||||
MochiRMSNormZero)
|
||||
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.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
@@ -54,8 +50,7 @@ class FeedForward(HF_FeedForward):
|
||||
inner_dim=None,
|
||||
bias: bool = True,
|
||||
):
|
||||
super().__init__(dim, dim_out, mult, dropout, activation_fn,
|
||||
final_dropout, inner_dim, bias)
|
||||
super().__init__(dim, dim_out, mult, dropout, activation_fn, final_dropout, inner_dim, bias)
|
||||
assert activation_fn == "swiglu"
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
@@ -100,26 +95,17 @@ class MochiAttention(nn.Module):
|
||||
self.to_k = nn.Linear(query_dim, self.inner_dim, bias=bias)
|
||||
self.to_v = nn.Linear(query_dim, self.inner_dim, bias=bias)
|
||||
|
||||
self.add_k_proj = nn.Linear(added_kv_proj_dim,
|
||||
self.inner_dim,
|
||||
bias=added_proj_bias)
|
||||
self.add_v_proj = nn.Linear(added_kv_proj_dim,
|
||||
self.inner_dim,
|
||||
bias=added_proj_bias)
|
||||
self.add_k_proj = nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias)
|
||||
self.add_v_proj = nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias)
|
||||
if self.context_pre_only is not None:
|
||||
self.add_q_proj = nn.Linear(added_kv_proj_dim,
|
||||
self.inner_dim,
|
||||
bias=added_proj_bias)
|
||||
self.add_q_proj = nn.Linear(added_kv_proj_dim, self.inner_dim, bias=added_proj_bias)
|
||||
|
||||
self.to_out = nn.ModuleList([])
|
||||
self.to_out.append(
|
||||
nn.Linear(self.inner_dim, self.out_dim, bias=out_bias))
|
||||
self.to_out.append(nn.Linear(self.inner_dim, self.out_dim, bias=out_bias))
|
||||
self.to_out.append(nn.Dropout(dropout))
|
||||
|
||||
if not self.context_pre_only:
|
||||
self.to_add_out = nn.Linear(self.inner_dim,
|
||||
self.out_context_dim,
|
||||
bias=out_bias)
|
||||
self.to_add_out = nn.Linear(self.inner_dim, self.out_context_dim, bias=out_bias)
|
||||
|
||||
self.processor = processor
|
||||
|
||||
@@ -144,9 +130,7 @@ class MochiAttnProcessor2_0:
|
||||
|
||||
def __init__(self):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError(
|
||||
"MochiAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0."
|
||||
)
|
||||
raise ImportError("MochiAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
@@ -198,9 +182,7 @@ class MochiAttnProcessor2_0:
|
||||
|
||||
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)
|
||||
@@ -241,11 +223,7 @@ class MochiAttnProcessor2_0:
|
||||
|
||||
attn_mask = encoder_attention_mask[:, :].bool()
|
||||
attn_mask = F.pad(attn_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 = flash_attn_no_pad(qkv, attn_mask, causal=False, dropout_p=0.0, softmax_scale=None)
|
||||
|
||||
# hidden_states = F.scaled_dot_product_attention(query, key, value, attn_mask = None, dropout_p=0.0, is_causal=False)
|
||||
|
||||
@@ -258,11 +236,8 @@ class MochiAttnProcessor2_0:
|
||||
hidden_states, encoder_hidden_states = hidden_states.split_with_sizes(
|
||||
(sequence_length, encoder_sequence_length), dim=1)
|
||||
# B, S, H, D
|
||||
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 = 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)
|
||||
@@ -324,16 +299,10 @@ class MochiTransformerBlock(nn.Module):
|
||||
self.ff_inner_dim = (4 * dim * 2) // 3
|
||||
self.ff_context_inner_dim = (4 * pooled_projection_dim * 2) // 3
|
||||
|
||||
self.norm1 = MochiRMSNormZero(dim,
|
||||
4 * dim,
|
||||
eps=eps,
|
||||
elementwise_affine=False)
|
||||
self.norm1 = MochiRMSNormZero(dim, 4 * dim, eps=eps, elementwise_affine=False)
|
||||
|
||||
if not context_pre_only:
|
||||
self.norm1_context = MochiRMSNormZero(dim,
|
||||
4 * pooled_projection_dim,
|
||||
eps=eps,
|
||||
elementwise_affine=False)
|
||||
self.norm1_context = MochiRMSNormZero(dim, 4 * pooled_projection_dim, eps=eps, elementwise_affine=False)
|
||||
else:
|
||||
self.norm1_context = MochiLayerNormContinuous(
|
||||
embedding_dim=pooled_projection_dim,
|
||||
@@ -357,17 +326,12 @@ class MochiTransformerBlock(nn.Module):
|
||||
|
||||
# TODO(aryan): norm_context layers are not needed when `context_pre_only` is True
|
||||
self.norm2 = MochiModulatedRMSNorm(eps=eps)
|
||||
self.norm2_context = (MochiModulatedRMSNorm(
|
||||
eps=eps) if not self.context_pre_only else None)
|
||||
self.norm2_context = (MochiModulatedRMSNorm(eps=eps) if not self.context_pre_only else None)
|
||||
|
||||
self.norm3 = MochiModulatedRMSNorm(eps)
|
||||
self.norm3_context = (MochiModulatedRMSNorm(
|
||||
eps=eps) if not self.context_pre_only else None)
|
||||
self.norm3_context = (MochiModulatedRMSNorm(eps=eps) if not self.context_pre_only else None)
|
||||
|
||||
self.ff = FeedForward(dim,
|
||||
inner_dim=self.ff_inner_dim,
|
||||
activation_fn=activation_fn,
|
||||
bias=False)
|
||||
self.ff = FeedForward(dim, inner_dim=self.ff_inner_dim, activation_fn=activation_fn, bias=False)
|
||||
self.ff_context = None
|
||||
if not context_pre_only:
|
||||
self.ff_context = FeedForward(
|
||||
@@ -389,8 +353,7 @@ class MochiTransformerBlock(nn.Module):
|
||||
image_rotary_emb: Optional[torch.Tensor] = None,
|
||||
output_attn=False,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
norm_hidden_states, gate_msa, scale_mlp, gate_mlp = self.norm1(
|
||||
hidden_states, temb)
|
||||
norm_hidden_states, gate_msa, scale_mlp, gate_mlp = self.norm1(hidden_states, temb)
|
||||
|
||||
if not self.context_pre_only:
|
||||
(
|
||||
@@ -400,8 +363,7 @@ class MochiTransformerBlock(nn.Module):
|
||||
enc_gate_mlp,
|
||||
) = self.norm1_context(encoder_hidden_states, temb)
|
||||
else:
|
||||
norm_encoder_hidden_states = self.norm1_context(
|
||||
encoder_hidden_states, temb)
|
||||
norm_encoder_hidden_states = self.norm1_context(encoder_hidden_states, temb)
|
||||
|
||||
attn_hidden_states, context_attn_hidden_states = self.attn1(
|
||||
hidden_states=norm_hidden_states,
|
||||
@@ -410,28 +372,21 @@ class MochiTransformerBlock(nn.Module):
|
||||
encoder_attention_mask=encoder_attention_mask,
|
||||
)
|
||||
|
||||
hidden_states = hidden_states + self.norm2(
|
||||
attn_hidden_states,
|
||||
torch.tanh(gate_msa).unsqueeze(1))
|
||||
norm_hidden_states = self.norm3(
|
||||
hidden_states, (1 + scale_mlp.unsqueeze(1).to(torch.float32)))
|
||||
hidden_states = hidden_states + self.norm2(attn_hidden_states, torch.tanh(gate_msa).unsqueeze(1))
|
||||
norm_hidden_states = self.norm3(hidden_states, (1 + scale_mlp.unsqueeze(1).to(torch.float32)))
|
||||
ff_output = self.ff(norm_hidden_states)
|
||||
hidden_states = hidden_states + self.norm4(
|
||||
ff_output,
|
||||
torch.tanh(gate_mlp).unsqueeze(1))
|
||||
hidden_states = hidden_states + self.norm4(ff_output, torch.tanh(gate_mlp).unsqueeze(1))
|
||||
|
||||
if not self.context_pre_only:
|
||||
encoder_hidden_states = encoder_hidden_states + self.norm2_context(
|
||||
context_attn_hidden_states,
|
||||
torch.tanh(enc_gate_msa).unsqueeze(1))
|
||||
encoder_hidden_states = encoder_hidden_states + self.norm2_context(context_attn_hidden_states,
|
||||
torch.tanh(enc_gate_msa).unsqueeze(1))
|
||||
norm_encoder_hidden_states = self.norm3_context(
|
||||
encoder_hidden_states,
|
||||
(1 + enc_scale_mlp.unsqueeze(1).to(torch.float32)),
|
||||
)
|
||||
context_ff_output = self.ff_context(norm_encoder_hidden_states)
|
||||
encoder_hidden_states = encoder_hidden_states + self.norm4_context(
|
||||
context_ff_output,
|
||||
torch.tanh(enc_gate_mlp).unsqueeze(1))
|
||||
encoder_hidden_states = encoder_hidden_states + self.norm4_context(context_ff_output,
|
||||
torch.tanh(enc_gate_mlp).unsqueeze(1))
|
||||
|
||||
if not output_attn:
|
||||
attn_hidden_states = None
|
||||
@@ -455,11 +410,7 @@ class MochiRoPE(nn.Module):
|
||||
self.target_area = base_height * base_width
|
||||
|
||||
def _centers(self, start, stop, num, device, dtype) -> torch.Tensor:
|
||||
edges = torch.linspace(start,
|
||||
stop,
|
||||
num + 1,
|
||||
device=device,
|
||||
dtype=dtype)
|
||||
edges = torch.linspace(start, stop, num + 1, device=device, dtype=dtype)
|
||||
return (edges[:-1] + edges[1:]) / 2
|
||||
|
||||
def _get_positions(
|
||||
@@ -471,21 +422,16 @@ class MochiRoPE(nn.Module):
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
) -> torch.Tensor:
|
||||
scale = (self.target_area / (height * width))**0.5
|
||||
t = torch.arange(num_frames * nccl_info.sp_size,
|
||||
device=device,
|
||||
dtype=dtype)
|
||||
h = self._centers(-height * scale / 2, height * scale / 2, height,
|
||||
device, dtype)
|
||||
w = self._centers(-width * scale / 2, width * scale / 2, width, device,
|
||||
dtype)
|
||||
t = torch.arange(num_frames * nccl_info.sp_size, device=device, dtype=dtype)
|
||||
h = self._centers(-height * scale / 2, height * scale / 2, height, device, dtype)
|
||||
w = self._centers(-width * scale / 2, width * scale / 2, width, device, dtype)
|
||||
|
||||
grid_t, grid_h, grid_w = torch.meshgrid(t, h, w, indexing="ij")
|
||||
|
||||
positions = torch.stack([grid_t, grid_h, grid_w], dim=-1).view(-1, 3)
|
||||
return positions
|
||||
|
||||
def _create_rope(self, freqs: torch.Tensor,
|
||||
pos: torch.Tensor) -> torch.Tensor:
|
||||
def _create_rope(self, freqs: torch.Tensor, pos: torch.Tensor) -> torch.Tensor:
|
||||
with torch.autocast(freqs.device.type, enabled=False):
|
||||
# Always run ROPE freqs computation in FP32
|
||||
freqs = torch.einsum(
|
||||
@@ -578,8 +524,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
num_attention_heads=8,
|
||||
)
|
||||
|
||||
self.pos_frequencies = nn.Parameter(
|
||||
torch.full((3, num_attention_heads, attention_head_dim // 2), 0.0))
|
||||
self.pos_frequencies = nn.Parameter(torch.full((3, num_attention_heads, attention_head_dim // 2), 0.0))
|
||||
self.rope = MochiRoPE()
|
||||
|
||||
self.transformer_blocks = nn.ModuleList([
|
||||
@@ -601,8 +546,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
eps=1e-6,
|
||||
norm_type="layer_norm",
|
||||
)
|
||||
self.proj_out = nn.Linear(inner_dim,
|
||||
patch_size * patch_size * out_channels)
|
||||
self.proj_out = nn.Linear(inner_dim, patch_size * patch_size * out_channels)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
@@ -621,8 +565,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
return_dict: bool = False,
|
||||
) -> torch.Tensor:
|
||||
assert (return_dict is False
|
||||
), "return_dict is not supported in MochiTransformer3DModel"
|
||||
assert (return_dict is False), "return_dict is not supported in MochiTransformer3DModel"
|
||||
|
||||
if attention_kwargs is not None:
|
||||
attention_kwargs = attention_kwargs.copy()
|
||||
@@ -634,11 +577,8 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
# 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."
|
||||
)
|
||||
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 = self.config.patch_size
|
||||
@@ -656,8 +596,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
|
||||
hidden_states = hidden_states.permute(0, 2, 1, 3, 4).flatten(0, 1)
|
||||
hidden_states = self.patch_embed(hidden_states)
|
||||
hidden_states = hidden_states.unflatten(0, (batch_size, -1)).flatten(
|
||||
1, 2)
|
||||
hidden_states = hidden_states.unflatten(0, (batch_size, -1)).flatten(1, 2)
|
||||
|
||||
image_rotary_emb = self.rope(
|
||||
self.pos_frequencies,
|
||||
@@ -678,9 +617,7 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
|
||||
return custom_forward
|
||||
|
||||
ckpt_kwargs: Dict[str, Any] = ({
|
||||
"use_reentrant": False
|
||||
} if is_torch_version(">=", "1.11.0") else {})
|
||||
ckpt_kwargs: Dict[str, Any] = ({"use_reentrant": False} if is_torch_version(">=", "1.11.0") else {})
|
||||
(
|
||||
hidden_states,
|
||||
encoder_hidden_states,
|
||||
@@ -710,12 +647,9 @@ class MochiTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin):
|
||||
hidden_states = self.norm_out(hidden_states, temb)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(batch_size, num_frames,
|
||||
post_patch_height,
|
||||
post_patch_width, p, p, -1)
|
||||
hidden_states = hidden_states.reshape(batch_size, num_frames, post_patch_height, post_patch_width, p, p, -1)
|
||||
hidden_states = hidden_states.permute(0, 6, 1, 2, 4, 3, 5)
|
||||
output = hidden_states.reshape(batch_size, -1, num_frames, height,
|
||||
width)
|
||||
output = hidden_states.reshape(batch_size, -1, num_frames, height, width)
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# remove `lora_scale` from each PEFT layer
|
||||
|
||||
@@ -78,9 +78,7 @@ class MochiLayerNormContinuous(nn.Module):
|
||||
|
||||
# AdaLN
|
||||
self.silu = nn.SiLU()
|
||||
self.linear_1 = nn.Linear(conditioning_embedding_dim,
|
||||
embedding_dim,
|
||||
bias=bias)
|
||||
self.linear_1 = nn.Linear(conditioning_embedding_dim, embedding_dim, bias=bias)
|
||||
self.norm = MochiModulatedRMSNorm(eps=eps)
|
||||
|
||||
def forward(
|
||||
@@ -117,16 +115,14 @@ class MochiRMSNormZero(nn.Module):
|
||||
self.linear = nn.Linear(embedding_dim, hidden_dim)
|
||||
self.norm = MochiModulatedRMSNorm(eps=eps)
|
||||
|
||||
def forward(
|
||||
self, hidden_states: torch.Tensor, emb: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
def forward(self, hidden_states: torch.Tensor,
|
||||
emb: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
hidden_states_dtype = hidden_states.dtype
|
||||
|
||||
emb = self.linear(self.silu(emb))
|
||||
scale_msa, gate_msa, scale_mlp, gate_mlp = emb.chunk(4, dim=1)
|
||||
|
||||
hidden_states = self.norm(hidden_states,
|
||||
(1 + scale_msa[:, None].to(torch.float32)))
|
||||
hidden_states = self.norm(hidden_states, (1 + scale_msa[:, None].to(torch.float32)))
|
||||
hidden_states = hidden_states.to(hidden_states_dtype)
|
||||
|
||||
return hidden_states, gate_msa, scale_mlp, gate_mlp
|
||||
|
||||
@@ -24,8 +24,7 @@ from diffusers.models.autoencoders import AutoencoderKL
|
||||
from diffusers.pipelines.mochi.pipeline_output import MochiPipelineOutput
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import (is_torch_xla_available, logging,
|
||||
replace_example_docstring)
|
||||
from diffusers.utils import is_torch_xla_available, logging, replace_example_docstring
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from einops import rearrange
|
||||
@@ -33,8 +32,7 @@ from transformers import T5EncoderModel, T5TokenizerFast
|
||||
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
|
||||
from fastvideo.utils.communications import all_gather
|
||||
from fastvideo.utils.parallel_states import (get_sequence_parallel_state,
|
||||
nccl_info)
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
|
||||
if is_torch_xla_available():
|
||||
import torch_xla.core.xla_model as xm
|
||||
@@ -78,19 +76,14 @@ def calculate_shift(
|
||||
def linear_quadratic_schedule(num_steps, threshold_noise, linear_steps=None):
|
||||
if linear_steps is None:
|
||||
linear_steps = num_steps // 2
|
||||
linear_sigma_schedule = [
|
||||
i * threshold_noise / linear_steps for i in range(linear_steps)
|
||||
]
|
||||
linear_sigma_schedule = [i * threshold_noise / linear_steps for i in range(linear_steps)]
|
||||
threshold_noise_step_diff = linear_steps - threshold_noise * num_steps
|
||||
quadratic_steps = num_steps - linear_steps
|
||||
quadratic_coef = threshold_noise_step_diff / (linear_steps *
|
||||
quadratic_steps**2)
|
||||
linear_coef = threshold_noise / linear_steps - 2 * threshold_noise_step_diff / (
|
||||
quadratic_steps**2)
|
||||
quadratic_coef = threshold_noise_step_diff / (linear_steps * quadratic_steps**2)
|
||||
linear_coef = threshold_noise / linear_steps - 2 * threshold_noise_step_diff / (quadratic_steps**2)
|
||||
const = quadratic_coef * (linear_steps**2)
|
||||
quadratic_sigma_schedule = [
|
||||
quadratic_coef * (i**2) + linear_coef * i + const
|
||||
for i in range(linear_steps, num_steps)
|
||||
quadratic_coef * (i**2) + linear_coef * i + const for i in range(linear_steps, num_steps)
|
||||
]
|
||||
sigma_schedule = linear_sigma_schedule + quadratic_sigma_schedule
|
||||
sigma_schedule = [1.0 - x for x in sigma_schedule]
|
||||
@@ -130,28 +123,22 @@ 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)
|
||||
@@ -187,9 +174,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
|
||||
model_cpu_offload_seq = "text_encoder->transformer->vae"
|
||||
_optional_components = []
|
||||
_callback_tensor_inputs = [
|
||||
"latents", "prompt_embeds", "negative_prompt_embeds"
|
||||
]
|
||||
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -212,11 +197,9 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
self.vae_temporal_scale_factor = 6
|
||||
self.patch_size = 2
|
||||
|
||||
self.video_processor = VideoProcessor(
|
||||
vae_scale_factor=self.vae_spatial_scale_factor)
|
||||
self.video_processor = VideoProcessor(vae_scale_factor=self.vae_spatial_scale_factor)
|
||||
self.tokenizer_max_length = (self.tokenizer.model_max_length
|
||||
if hasattr(self, "tokenizer")
|
||||
and self.tokenizer is not None else 77)
|
||||
if hasattr(self, "tokenizer") and self.tokenizer is not None else 77)
|
||||
self.default_height = 480
|
||||
self.default_width = 848
|
||||
|
||||
@@ -247,31 +230,23 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
prompt_attention_mask = text_inputs.attention_mask
|
||||
prompt_attention_mask = prompt_attention_mask.bool().to(device)
|
||||
|
||||
untruncated_ids = self.tokenizer(prompt,
|
||||
padding="longest",
|
||||
return_tensors="pt").input_ids
|
||||
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
|
||||
|
||||
if untruncated_ids.shape[-1] >= text_input_ids.shape[
|
||||
-1] and not torch.equal(text_input_ids, untruncated_ids):
|
||||
removed_text = self.tokenizer.batch_decode(
|
||||
untruncated_ids[:, max_sequence_length - 1:-1])
|
||||
logger.warning(
|
||||
"The following part of your input was truncated because `max_sequence_length` is set to "
|
||||
f" {max_sequence_length} tokens: {removed_text}")
|
||||
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
|
||||
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, max_sequence_length - 1:-1])
|
||||
logger.warning("The following part of your input was truncated because `max_sequence_length` is set to "
|
||||
f" {max_sequence_length} tokens: {removed_text}")
|
||||
|
||||
prompt_embeds = self.text_encoder(
|
||||
text_input_ids.to(device), attention_mask=prompt_attention_mask)[0]
|
||||
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=prompt_attention_mask)[0]
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
_, seq_len, _ = prompt_embeds.shape
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
|
||||
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt,
|
||||
seq_len, -1)
|
||||
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1)
|
||||
|
||||
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
|
||||
prompt_attention_mask = prompt_attention_mask.repeat(
|
||||
num_videos_per_prompt, 1)
|
||||
prompt_attention_mask = prompt_attention_mask.repeat(num_videos_per_prompt, 1)
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
@@ -335,11 +310,9 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
|
||||
if do_classifier_free_guidance and negative_prompt_embeds is None:
|
||||
negative_prompt = negative_prompt or ""
|
||||
negative_prompt = (batch_size * [negative_prompt] if isinstance(
|
||||
negative_prompt, str) else negative_prompt)
|
||||
negative_prompt = (batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt)
|
||||
|
||||
if prompt is not None and type(prompt) is not type(
|
||||
negative_prompt):
|
||||
if prompt is not None and type(prompt) is not type(negative_prompt):
|
||||
raise TypeError(
|
||||
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
|
||||
f" {type(prompt)}.")
|
||||
@@ -379,13 +352,10 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
negative_prompt_attention_mask=None,
|
||||
):
|
||||
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 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_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,24 +366,15 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
" 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 prompt_embeds is not None and prompt_attention_mask is None:
|
||||
raise ValueError(
|
||||
"Must provide `prompt_attention_mask` when specifying `prompt_embeds`."
|
||||
)
|
||||
raise ValueError("Must provide `prompt_attention_mask` when specifying `prompt_embeds`.")
|
||||
|
||||
if (negative_prompt_embeds is not None
|
||||
and negative_prompt_attention_mask is None):
|
||||
raise ValueError(
|
||||
"Must provide `negative_prompt_attention_mask` when specifying `negative_prompt_embeds`."
|
||||
)
|
||||
if (negative_prompt_embeds is not None and negative_prompt_attention_mask is None):
|
||||
raise ValueError("Must provide `negative_prompt_attention_mask` when specifying `negative_prompt_embeds`.")
|
||||
|
||||
if prompt_embeds is not None and negative_prompt_embeds is not None:
|
||||
if prompt_embeds.shape != negative_prompt_embeds.shape:
|
||||
@@ -479,13 +440,9 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
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.")
|
||||
|
||||
latents = randn_tensor(shape,
|
||||
generator=generator,
|
||||
device=device,
|
||||
dtype=torch.float32)
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=torch.float32)
|
||||
latents = latents.to(dtype)
|
||||
return latents
|
||||
|
||||
@@ -522,8 +479,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
timesteps: List[int] = None,
|
||||
guidance_scale: float = 4.5,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator,
|
||||
List[torch.Generator]]] = None,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
prompt_attention_mask: Optional[torch.Tensor] = None,
|
||||
@@ -532,8 +488,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
output_type: Optional[str] = "pil",
|
||||
return_dict: bool = True,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
callback_on_step_end: Optional[Callable[[int, int, Dict],
|
||||
None]] = None,
|
||||
callback_on_step_end: Optional[Callable[[int, int, Dict], None]] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
max_sequence_length: int = 256,
|
||||
return_all_states=False,
|
||||
@@ -612,8 +567,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
is returned where the first element is a list with the generated images.
|
||||
"""
|
||||
|
||||
if isinstance(callback_on_step_end,
|
||||
(PipelineCallback, MultiPipelineCallbacks)):
|
||||
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
|
||||
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
|
||||
|
||||
height = height or self.default_height
|
||||
@@ -624,8 +578,7 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
prompt=prompt,
|
||||
height=height,
|
||||
width=width,
|
||||
callback_on_step_end_tensor_inputs=
|
||||
callback_on_step_end_tensor_inputs,
|
||||
callback_on_step_end_tensor_inputs=callback_on_step_end_tensor_inputs,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
prompt_attention_mask=prompt_attention_mask,
|
||||
@@ -665,10 +618,8 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
device=device,
|
||||
)
|
||||
if self.do_classifier_free_guidance:
|
||||
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds],
|
||||
dim=0)
|
||||
prompt_attention_mask = torch.cat(
|
||||
[negative_prompt_attention_mask, prompt_attention_mask], dim=0)
|
||||
prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds], dim=0)
|
||||
prompt_attention_mask = torch.cat([negative_prompt_attention_mask, prompt_attention_mask], dim=0)
|
||||
|
||||
# 4. Prepare latent variables
|
||||
num_channels_latents = self.transformer.config.in_channels
|
||||
@@ -685,17 +636,14 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
)
|
||||
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, :, :, :]
|
||||
|
||||
original_noise = copy.deepcopy(latents)
|
||||
# 5. Prepare timestep
|
||||
# from https://github.com/genmoai/models/blob/075b6e36db58f1242921deff83a1066887b9c9e1/src/mochi_preview/infer.py#L77
|
||||
threshold_noise = 0.025
|
||||
sigmas = linear_quadratic_schedule(num_inference_steps,
|
||||
threshold_noise)
|
||||
sigmas = linear_quadratic_schedule(num_inference_steps, threshold_noise)
|
||||
sigmas = np.array(sigmas)
|
||||
# check if of type FlowMatchEulerDiscreteScheduler
|
||||
if isinstance(self.scheduler, FlowMatchEulerDiscreteScheduler):
|
||||
@@ -712,25 +660,19 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
num_inference_steps,
|
||||
device,
|
||||
)
|
||||
num_warmup_steps = max(
|
||||
len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
num_warmup_steps = max(len(timesteps) - num_inference_steps * self.scheduler.order, 0)
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
# 6. Denoising loop
|
||||
self._progress_bar_config = {
|
||||
"disable": nccl_info.rank_within_group != 0
|
||||
}
|
||||
self._progress_bar_config = {"disable": nccl_info.rank_within_group != 0}
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
latent_model_input = (torch.cat(
|
||||
[latents] *
|
||||
2) if self.do_classifier_free_guidance else latents)
|
||||
latent_model_input = (torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents)
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latent_model_input.shape[0]).to(
|
||||
latents.dtype)
|
||||
timestep = t.expand(latent_model_input.shape[0]).to(latents.dtype)
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
@@ -745,15 +687,11 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
noise_pred = noise_pred.to(torch.float32)
|
||||
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)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents_dtype = latents.dtype
|
||||
latents = self.scheduler.step(noise_pred,
|
||||
t,
|
||||
latents.to(torch.float32),
|
||||
return_dict=False)[0]
|
||||
latents = self.scheduler.step(noise_pred, t, latents.to(torch.float32), return_dict=False)[0]
|
||||
latents = latents.to(latents_dtype)
|
||||
|
||||
if latents.dtype != latents_dtype:
|
||||
@@ -765,17 +703,13 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
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)
|
||||
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)
|
||||
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):
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
|
||||
if XLA_AVAILABLE:
|
||||
@@ -795,25 +729,19 @@ class MochiPipeline(DiffusionPipeline, Mochi1LoraLoaderMixin):
|
||||
else:
|
||||
# unscale/denormalize the latents
|
||||
# denormalize with the mean and std if available and not None
|
||||
has_latents_mean = (hasattr(self.vae.config, "latents_mean")
|
||||
and self.vae.config.latents_mean is not None)
|
||||
has_latents_std = (hasattr(self.vae.config, "latents_std")
|
||||
and self.vae.config.latents_std is not None)
|
||||
has_latents_mean = (hasattr(self.vae.config, "latents_mean") and self.vae.config.latents_mean is not None)
|
||||
has_latents_std = (hasattr(self.vae.config, "latents_std") and self.vae.config.latents_std is not None)
|
||||
if has_latents_mean and has_latents_std:
|
||||
latents_mean = (torch.tensor(
|
||||
self.vae.config.latents_mean).view(1, 12, 1, 1, 1).to(
|
||||
latents.device, latents.dtype))
|
||||
latents_std = (torch.tensor(self.vae.config.latents_std).view(
|
||||
1, 12, 1, 1, 1).to(latents.device, latents.dtype))
|
||||
latents = (
|
||||
latents * latents_std / self.vae.config.scaling_factor +
|
||||
latents_mean)
|
||||
latents_mean = (torch.tensor(self.vae.config.latents_mean).view(1, 12, 1, 1,
|
||||
1).to(latents.device, latents.dtype))
|
||||
latents_std = (torch.tensor(self.vae.config.latents_std).view(1, 12, 1, 1,
|
||||
1).to(latents.device, latents.dtype))
|
||||
latents = (latents * latents_std / self.vae.config.scaling_factor + latents_mean)
|
||||
else:
|
||||
latents = latents / 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)
|
||||
video = self.video_processor.postprocess_video(video, output_type=output_type)
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
@@ -0,0 +1,7 @@
|
||||
import os
|
||||
|
||||
os.environ["NCCL_DEBUG"] = "ERROR"
|
||||
|
||||
from .diffusion.scheduler import *
|
||||
from .diffusion.video_pipeline import *
|
||||
from .modules.model import *
|
||||
@@ -0,0 +1 @@
|
||||
__version__ = "0.1.0"
|
||||
@@ -0,0 +1,174 @@
|
||||
import argparse
|
||||
|
||||
|
||||
def parse_args(namespace=None):
|
||||
parser = argparse.ArgumentParser(description="StepVideo inference script")
|
||||
|
||||
parser = add_extra_models_args(parser)
|
||||
parser = add_denoise_schedule_args(parser)
|
||||
parser = add_inference_args(parser)
|
||||
parser = add_parallel_args(parser)
|
||||
|
||||
args = parser.parse_args(namespace=namespace)
|
||||
|
||||
return args
|
||||
|
||||
|
||||
def add_extra_models_args(parser: argparse.ArgumentParser):
|
||||
group = parser.add_argument_group(title="Extra models args, including vae, text encoders and tokenizers)")
|
||||
|
||||
group.add_argument(
|
||||
"--vae_url",
|
||||
type=str,
|
||||
default='127.0.0.1',
|
||||
help="vae url.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--caption_url",
|
||||
type=str,
|
||||
default='127.0.0.1',
|
||||
help="caption url.",
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
def add_denoise_schedule_args(parser: argparse.ArgumentParser):
|
||||
group = parser.add_argument_group(title="Denoise schedule args")
|
||||
|
||||
# Flow Matching
|
||||
group.add_argument(
|
||||
"--time_shift",
|
||||
type=float,
|
||||
default=7.0,
|
||||
help="Shift factor for flow matching schedulers.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--flow_reverse",
|
||||
action="store_true",
|
||||
help="If reverse, learning/sampling from t=1 -> t=0.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--flow_solver",
|
||||
type=str,
|
||||
default="euler",
|
||||
help="Solver for flow matching.",
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
def add_inference_args(parser: argparse.ArgumentParser):
|
||||
group = parser.add_argument_group(title="Inference args")
|
||||
|
||||
# ======================== Model loads ========================
|
||||
group.add_argument(
|
||||
"--model_dir",
|
||||
type=str,
|
||||
default="./ckpts",
|
||||
help="Root path of all the models, including t2v models and extra models.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--model_resolution",
|
||||
type=str,
|
||||
default="540p",
|
||||
choices=["540p"],
|
||||
help="Root path of all the models, including t2v models and extra models.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--use-cpu-offload",
|
||||
action="store_true",
|
||||
help="Use CPU offload for the model load.",
|
||||
)
|
||||
|
||||
# ======================== Inference general setting ========================
|
||||
group.add_argument(
|
||||
"--batch_size",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Batch size for inference and evaluation.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--infer_steps",
|
||||
type=int,
|
||||
default=50,
|
||||
help="Number of denoising steps for inference.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--save_path",
|
||||
type=str,
|
||||
default="./results",
|
||||
help="Path to save the generated samples.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--name_suffix",
|
||||
type=str,
|
||||
default="",
|
||||
help="Suffix for the names of saved samples.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--num_videos",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of videos to generate for each prompt.",
|
||||
)
|
||||
# ---sample size---
|
||||
group.add_argument(
|
||||
"--num_frames",
|
||||
type=int,
|
||||
default=204,
|
||||
help="How many frames to sample from a video. ",
|
||||
)
|
||||
group.add_argument(
|
||||
"--height",
|
||||
type=int,
|
||||
default=544,
|
||||
help="The height of video sample",
|
||||
)
|
||||
group.add_argument(
|
||||
"--width",
|
||||
type=int,
|
||||
default=992,
|
||||
help="The width of video sample",
|
||||
)
|
||||
# --- prompt ---
|
||||
group.add_argument(
|
||||
"--prompt",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Prompt for sampling during evaluation.",
|
||||
)
|
||||
group.add_argument("--seed", type=int, default=1234, help="Seed for evaluation.")
|
||||
|
||||
# Classifier-Free Guidance
|
||||
group.add_argument("--pos_magic",
|
||||
type=str,
|
||||
default="超高清、HDR 视频、环境光、杜比全景声、画面稳定、流畅动作、逼真的细节、专业级构图、超现实主义、自然、生动、超细节、清晰。",
|
||||
help="Positive magic prompt for sampling.")
|
||||
group.add_argument("--neg_magic",
|
||||
type=str,
|
||||
default="画面暗、低分辨率、不良手、文本、缺少手指、多余的手指、裁剪、低质量、颗粒状、签名、水印、用户名、模糊。",
|
||||
help="Negative magic prompt for sampling.")
|
||||
group.add_argument("--cfg_scale", type=float, default=9.0, help="Classifier free guidance scale.")
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
def add_parallel_args(parser: argparse.ArgumentParser):
|
||||
group = parser.add_argument_group(title="Parallel args")
|
||||
|
||||
# ======================== Model loads ========================
|
||||
group.add_argument(
|
||||
"--ulysses_degree",
|
||||
type=int,
|
||||
default=8,
|
||||
help="Ulysses degree.",
|
||||
)
|
||||
group.add_argument(
|
||||
"--ring_degree",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Ulysses degree.",
|
||||
)
|
||||
|
||||
return parser
|
||||
@@ -0,0 +1,220 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.schedulers.scheduling_utils import SchedulerMixin
|
||||
from diffusers.utils import BaseOutput, logging
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
@dataclass
|
||||
class FlowMatchDiscreteSchedulerOutput(BaseOutput):
|
||||
"""
|
||||
Output class for the scheduler's `step` function output.
|
||||
|
||||
Args:
|
||||
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
|
||||
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
|
||||
denoising loop.
|
||||
"""
|
||||
|
||||
prev_sample: torch.FloatTensor
|
||||
|
||||
|
||||
class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
"""
|
||||
Euler scheduler.
|
||||
|
||||
This model inherits from [`SchedulerMixin`] and [`ConfigMixin`]. Check the superclass documentation for the generic
|
||||
methods the library implements for all schedulers such as loading and saving.
|
||||
|
||||
Args:
|
||||
num_train_timesteps (`int`, defaults to 1000):
|
||||
The number of diffusion steps to train the model.
|
||||
timestep_spacing (`str`, defaults to `"linspace"`):
|
||||
The way the timesteps should be scaled. Refer to Table 2 of the [Common Diffusion Noise Schedules and
|
||||
Sample Steps are Flawed](https://huggingface.co/papers/2305.08891) for more information.
|
||||
reverse (`bool`, defaults to `True`):
|
||||
Whether to reverse the timestep schedule.
|
||||
"""
|
||||
|
||||
_compatibles = []
|
||||
order = 1
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
num_train_timesteps: int = 1000,
|
||||
reverse: bool = False,
|
||||
solver: str = "euler",
|
||||
device: Union[str, torch.device] = None,
|
||||
):
|
||||
sigmas = torch.linspace(1, 0, num_train_timesteps + 1)
|
||||
|
||||
if not reverse:
|
||||
sigmas = sigmas.flip(0)
|
||||
|
||||
self.sigmas = sigmas
|
||||
# the value fed to model
|
||||
self.timesteps = (sigmas[:-1] * num_train_timesteps).to(dtype=torch.float32)
|
||||
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
self.device = device
|
||||
|
||||
self.supported_solver = ["euler"]
|
||||
if solver not in self.supported_solver:
|
||||
raise ValueError(f"Solver {solver} not supported. Supported solvers: {self.supported_solver}")
|
||||
|
||||
@property
|
||||
def step_index(self):
|
||||
"""
|
||||
The index counter for current timestep. It will increase 1 after each scheduler step.
|
||||
"""
|
||||
return self._step_index
|
||||
|
||||
@property
|
||||
def begin_index(self):
|
||||
"""
|
||||
The index for the first timestep. It should be set from pipeline with `set_begin_index` method.
|
||||
"""
|
||||
return self._begin_index
|
||||
|
||||
# Copied from diffusers.schedulers.scheduling_dpmsolver_multistep.DPMSolverMultistepScheduler.set_begin_index
|
||||
def set_begin_index(self, begin_index: int = 0):
|
||||
"""
|
||||
Sets the begin index for the scheduler. This function should be run from pipeline before the inference.
|
||||
|
||||
Args:
|
||||
begin_index (`int`):
|
||||
The begin index for the scheduler.
|
||||
"""
|
||||
self._begin_index = begin_index
|
||||
|
||||
def _sigma_to_t(self, sigma):
|
||||
return sigma * self.config.num_train_timesteps
|
||||
|
||||
def set_timesteps(
|
||||
self,
|
||||
num_inference_steps: int,
|
||||
time_shift: float = 13.0,
|
||||
device: Union[str, torch.device] = None,
|
||||
):
|
||||
"""
|
||||
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
|
||||
Args:
|
||||
num_inference_steps (`int`):
|
||||
The number of diffusion steps used when generating samples with a pre-trained model.
|
||||
device (`str` or `torch.device`, *optional*):
|
||||
The device to which the timesteps should be moved to. If `None`, the timesteps are not moved.
|
||||
n_tokens (`int`, *optional*):
|
||||
Number of tokens in the input sequence.
|
||||
"""
|
||||
device = device or self.device
|
||||
self.num_inference_steps = num_inference_steps
|
||||
|
||||
sigmas = torch.linspace(1, 0, num_inference_steps + 1, device=device)
|
||||
sigmas = self.sd3_time_shift(sigmas, time_shift)
|
||||
|
||||
if not self.config.reverse:
|
||||
sigmas = 1 - sigmas
|
||||
|
||||
self.sigmas = sigmas
|
||||
self.timesteps = sigmas[:-1]
|
||||
|
||||
# Reset step index
|
||||
self._step_index = None
|
||||
|
||||
def index_for_timestep(self, timestep, schedule_timesteps=None):
|
||||
if schedule_timesteps is None:
|
||||
schedule_timesteps = self.timesteps
|
||||
|
||||
indices = (schedule_timesteps == timestep).nonzero()
|
||||
|
||||
# The sigma index that is taken for the **very** first `step`
|
||||
# is always the second index (or the last index if there is only 1)
|
||||
# This way we can ensure we don't accidentally skip a sigma in
|
||||
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
|
||||
pos = 1 if len(indices) > 1 else 0
|
||||
|
||||
return indices[pos].item()
|
||||
|
||||
def _init_step_index(self, timestep):
|
||||
if self.begin_index is None:
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
timestep = timestep.to(self.timesteps.device)
|
||||
self._step_index = self.index_for_timestep(timestep)
|
||||
else:
|
||||
self._step_index = self._begin_index
|
||||
|
||||
def scale_model_input(self, sample: torch.Tensor, timestep: Optional[int] = None) -> torch.Tensor:
|
||||
return sample
|
||||
|
||||
def sd3_time_shift(self, t: torch.Tensor, time_shift: float = 13.0):
|
||||
return (time_shift * t) / (1 + (time_shift - 1) * t)
|
||||
|
||||
def step(
|
||||
self,
|
||||
model_output: torch.FloatTensor,
|
||||
timestep: Union[float, torch.FloatTensor],
|
||||
sample: torch.FloatTensor,
|
||||
return_dict: bool = False,
|
||||
) -> Union[FlowMatchDiscreteSchedulerOutput, Tuple]:
|
||||
"""
|
||||
Predict the sample from the previous timestep by reversing the SDE. This function propagates the diffusion
|
||||
process from the learned model outputs (most often the predicted noise).
|
||||
|
||||
Args:
|
||||
model_output (`torch.FloatTensor`):
|
||||
The direct output from learned diffusion model.
|
||||
timestep (`float`):
|
||||
The current discrete timestep in the diffusion chain.
|
||||
sample (`torch.FloatTensor`):
|
||||
A current instance of a sample created by the diffusion process.
|
||||
generator (`torch.Generator`, *optional*):
|
||||
A random number generator.
|
||||
n_tokens (`int`, *optional*):
|
||||
Number of tokens in the input sequence.
|
||||
return_dict (`bool`):
|
||||
Whether or not to return a [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or
|
||||
tuple.
|
||||
|
||||
Returns:
|
||||
[`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] or `tuple`:
|
||||
If return_dict is `True`, [`~schedulers.scheduling_euler_discrete.EulerDiscreteSchedulerOutput`] is
|
||||
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 self.step_index is None:
|
||||
self._init_step_index(timestep)
|
||||
|
||||
# Upcast to avoid precision issues when computing prev_sample
|
||||
sample = sample.to(torch.float32)
|
||||
|
||||
dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
|
||||
|
||||
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}")
|
||||
|
||||
# upon completion increase step index by one
|
||||
self._step_index += 1
|
||||
|
||||
if not return_dict:
|
||||
return prev_sample
|
||||
|
||||
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
def __len__(self):
|
||||
return self.config.num_train_timesteps
|
||||
+325
@@ -0,0 +1,325 @@
|
||||
# Copyright 2025 StepFun Inc. All Rights Reserved.
|
||||
|
||||
import asyncio
|
||||
import pickle
|
||||
from dataclasses import dataclass
|
||||
from typing import Dict, List, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.utils import BaseOutput
|
||||
|
||||
from fastvideo.models.stepvideo.diffusion.scheduler import FlowMatchDiscreteScheduler
|
||||
from fastvideo.models.stepvideo.modules.model import StepVideoModel
|
||||
from fastvideo.models.stepvideo.utils import VideoProcessor
|
||||
|
||||
|
||||
def call_api_gen(url, api, port=8080):
|
||||
url = f"http://{url}:{port}/{api}-api"
|
||||
import aiohttp
|
||||
|
||||
async def _fn(samples, *args, **kwargs):
|
||||
if api == 'vae':
|
||||
data = {
|
||||
"samples": samples,
|
||||
}
|
||||
elif api == 'caption':
|
||||
data = {
|
||||
"prompts": samples,
|
||||
}
|
||||
else:
|
||||
raise Exception(f"Not supported api: {api}...")
|
||||
|
||||
async with aiohttp.ClientSession() as sess:
|
||||
data_bytes = pickle.dumps(data)
|
||||
async with sess.get(url, data=data_bytes, timeout=12000) as response:
|
||||
result = bytearray()
|
||||
while not response.content.at_eof():
|
||||
chunk = await response.content.read(1024)
|
||||
result += chunk
|
||||
response_data = pickle.loads(result)
|
||||
return response_data
|
||||
|
||||
return _fn
|
||||
|
||||
|
||||
@dataclass
|
||||
class StepVideoPipelineOutput(BaseOutput):
|
||||
video: Union[torch.Tensor, np.ndarray]
|
||||
|
||||
|
||||
class StepVideoPipeline(DiffusionPipeline):
|
||||
r"""
|
||||
Pipeline for text-to-video generation using StepVideo.
|
||||
|
||||
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:
|
||||
transformer ([`StepVideoModel`]):
|
||||
Conditional Transformer to denoise the encoded image latents.
|
||||
scheduler ([`FlowMatchDiscreteScheduler`]):
|
||||
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
|
||||
vae_url:
|
||||
remote vae server's url.
|
||||
caption_url:
|
||||
remote caption (stepllm and clip) server's url.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
transformer: StepVideoModel,
|
||||
scheduler: FlowMatchDiscreteScheduler,
|
||||
vae_url: str = '127.0.0.1',
|
||||
caption_url: str = '127.0.0.1',
|
||||
save_path: str = './results',
|
||||
name_suffix: str = '',
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
transformer=transformer,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
|
||||
self.vae_scale_factor_temporal = self.vae.temporal_compression_ratio if getattr(self, "vae", None) else 8
|
||||
self.vae_scale_factor_spatial = self.vae.spatial_compression_ratio if getattr(self, "vae", None) else 16
|
||||
self.video_processor = VideoProcessor(save_path, name_suffix)
|
||||
|
||||
self.vae_url = vae_url
|
||||
self.caption_url = caption_url
|
||||
self.setup_api(self.vae_url, self.caption_url)
|
||||
|
||||
def setup_api(self, vae_url, caption_url):
|
||||
self.vae_url = vae_url
|
||||
self.caption_url = caption_url
|
||||
self.caption = call_api_gen(caption_url, 'caption')
|
||||
self.vae = call_api_gen(vae_url, 'vae')
|
||||
return self
|
||||
|
||||
def encode_prompt(
|
||||
self,
|
||||
prompt: str,
|
||||
neg_magic: str = '',
|
||||
pos_magic: str = '',
|
||||
):
|
||||
device = self._execution_device
|
||||
prompts = [prompt + pos_magic]
|
||||
bs = len(prompts)
|
||||
prompts += [neg_magic] * bs
|
||||
|
||||
data = asyncio.run(self.caption(prompts))
|
||||
prompt_embeds, prompt_attention_mask, clip_embedding = data['y'].to(device), data['y_mask'].to(
|
||||
device), data['clip_embedding'].to(device)
|
||||
|
||||
return prompt_embeds, clip_embedding, prompt_attention_mask
|
||||
|
||||
def decode_vae(self, samples):
|
||||
samples = asyncio.run(self.vae(samples.cpu()))
|
||||
return samples
|
||||
|
||||
def check_inputs(self, num_frames, width, height):
|
||||
num_frames = max(num_frames // 17 * 17, 1)
|
||||
width = max(width // 16 * 16, 16)
|
||||
height = max(height // 16 * 16, 16)
|
||||
return num_frames, width, height
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
batch_size: int,
|
||||
num_channels_latents: 64,
|
||||
height: int = 544,
|
||||
width: int = 992,
|
||||
num_frames: int = 204,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if latents is not None:
|
||||
return latents.to(device=device, dtype=dtype)
|
||||
|
||||
num_frames, width, height = self.check_inputs(num_frames, width, height)
|
||||
shape = (
|
||||
batch_size,
|
||||
max(num_frames // 17 * 3, 1),
|
||||
num_channels_latents,
|
||||
int(height) // self.vae_scale_factor_spatial,
|
||||
int(width) // self.vae_scale_factor_spatial,
|
||||
) # b,f,c,h,w
|
||||
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.")
|
||||
|
||||
if generator is None:
|
||||
generator = torch.Generator(device=self._execution_device)
|
||||
|
||||
latents = torch.randn(shape, generator=generator, device=device, dtype=dtype)
|
||||
return latents
|
||||
|
||||
@torch.inference_mode()
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
height: int = 544,
|
||||
width: int = 992,
|
||||
num_frames: int = 204,
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 9.0,
|
||||
time_shift: float = 13.0,
|
||||
neg_magic: str = "",
|
||||
pos_magic: str = "",
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
output_type: Optional[str] = "mp4",
|
||||
output_file_name: Optional[str] = "",
|
||||
return_dict: bool = True,
|
||||
mask_strategy: Optional[Dict[str, list]] = None,
|
||||
):
|
||||
r"""
|
||||
The call function to the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||
instead.
|
||||
height (`int`, defaults to `544`):
|
||||
The height in pixels of the generated image.
|
||||
width (`int`, defaults to `992`):
|
||||
The width in pixels of the generated image.
|
||||
num_frames (`int`, defaults to `204`):
|
||||
The number of frames in the generated video.
|
||||
num_inference_steps (`int`, defaults to `50`):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
guidance_scale (`float`, defaults to `9.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.
|
||||
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`.
|
||||
output_type (`str`, *optional*, defaults to `"pil"`):
|
||||
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
|
||||
output_file_name(`str`, *optional*`):
|
||||
The output mp4 file name.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`StepVideoPipelineOutput`] instead of a plain tuple.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~StepVideoPipelineOutput`] or `tuple`:
|
||||
If `return_dict` is `True`, [`StepVideoPipelineOutput`] 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.
|
||||
"""
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
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]
|
||||
|
||||
do_classifier_free_guidance = guidance_scale > 1.0
|
||||
|
||||
# 3. Encode input prompt
|
||||
prompt_embeds, prompt_embeds_2, prompt_attention_mask = self.encode_prompt(
|
||||
prompt=prompt,
|
||||
neg_magic=neg_magic,
|
||||
pos_magic=pos_magic,
|
||||
)
|
||||
|
||||
transformer_dtype = self.transformer.dtype
|
||||
prompt_embeds = prompt_embeds.to(transformer_dtype)
|
||||
prompt_attention_mask = prompt_attention_mask.to(transformer_dtype)
|
||||
prompt_embeds_2 = prompt_embeds_2.to(transformer_dtype)
|
||||
|
||||
# 4. Prepare timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps=num_inference_steps, time_shift=time_shift, device=device)
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = self.transformer.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
num_frames,
|
||||
torch.bfloat16,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
def dict_to_3d_list(best_masks, t_max=50, l_max=48, h_max=48):
|
||||
result = [[[None for _ in range(h_max)] for _ in range(l_max)] for _ in range(t_max)]
|
||||
if best_masks is None:
|
||||
return result
|
||||
for key, value in best_masks.items():
|
||||
timestep, layer, head = map(int, key.split('_'))
|
||||
result[timestep][layer][head] = value
|
||||
return result
|
||||
|
||||
mask_strategy = dict_to_3d_list(mask_strategy)
|
||||
|
||||
#best_mask_selections = None
|
||||
# 7. Denoising loop
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(self.scheduler.timesteps):
|
||||
latent_model_input = torch.cat([latents] * 2) if do_classifier_free_guidance else latents
|
||||
latent_model_input = latent_model_input.to(transformer_dtype)
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latent_model_input.shape[0]).to(latent_model_input.dtype)
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
encoder_hidden_states_2=prompt_embeds_2,
|
||||
return_dict=False,
|
||||
mask_strategy=mask_strategy[i],
|
||||
)
|
||||
# perform guidance
|
||||
if do_classifier_free_guidance:
|
||||
noise_pred_text, noise_pred_uncond = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (noise_pred_text - noise_pred_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(model_output=noise_pred, timestep=t, sample=latents)
|
||||
|
||||
progress_bar.update()
|
||||
|
||||
if not torch.distributed.is_initialized() or int(torch.distributed.get_rank()) == 0:
|
||||
if not output_type == "latent":
|
||||
video = self.decode_vae(latents)
|
||||
video = self.video_processor.postprocess_video(video,
|
||||
output_file_name=output_file_name,
|
||||
output_type=output_type)
|
||||
else:
|
||||
video = latents
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return (video, )
|
||||
|
||||
return StepVideoPipelineOutput(video=video)
|
||||
+96
@@ -0,0 +1,96 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
from flash_attn import flash_attn_func
|
||||
|
||||
try:
|
||||
from st_attn import sliding_tile_attention
|
||||
except ImportError:
|
||||
print("Could not load Sliding Tile Attention.")
|
||||
sliding_tile_attention = None
|
||||
|
||||
from fastvideo.utils.communications import all_to_all_4D
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
|
||||
|
||||
class Attention(nn.Module):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def attn_processor(self, attn_type):
|
||||
if attn_type == 'torch':
|
||||
return self.torch_attn_func
|
||||
elif attn_type == 'parallel':
|
||||
return self.parallel_attn_func
|
||||
else:
|
||||
raise Exception('Not supported attention type...')
|
||||
|
||||
def tile(self, x, sp_size):
|
||||
x = rearrange(x, "b (sp t h w) head d -> b (t sp h w) head d", sp=sp_size, t=36 // sp_size, h=48, w=48)
|
||||
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=6,
|
||||
n_h=6,
|
||||
n_w=6,
|
||||
ts_t=6,
|
||||
ts_h=8,
|
||||
ts_w=8)
|
||||
|
||||
def untile(self, 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=6,
|
||||
n_h=6,
|
||||
n_w=6,
|
||||
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=36 // sp_size, h=48, w=48)
|
||||
|
||||
def torch_attn_func(self, q, k, v, attn_mask=None, causal=False, drop_rate=0.0, **kwargs):
|
||||
|
||||
if attn_mask is not None and attn_mask.dtype != torch.bool:
|
||||
attn_mask = attn_mask.to(q.dtype)
|
||||
|
||||
if attn_mask is not None and attn_mask.ndim == 3: ## no head
|
||||
n_heads = q.shape[2]
|
||||
attn_mask = attn_mask.unsqueeze(1).repeat(1, n_heads, 1, 1)
|
||||
|
||||
q, k, v = map(lambda x: rearrange(x, 'b s h d -> b h s d'), (q, k, v))
|
||||
x = torch.nn.functional.scaled_dot_product_attention(q,
|
||||
k,
|
||||
v,
|
||||
attn_mask=attn_mask,
|
||||
dropout_p=drop_rate,
|
||||
is_causal=causal)
|
||||
x = rearrange(x, 'b h s d -> b s h d')
|
||||
return x
|
||||
|
||||
def parallel_attn_func(self, q, k, v, causal=False, mask_strategy=None, **kwargs):
|
||||
if get_sequence_parallel_state():
|
||||
q = all_to_all_4D(q, scatter_dim=2, gather_dim=1)
|
||||
k = all_to_all_4D(k, scatter_dim=2, gather_dim=1)
|
||||
v = all_to_all_4D(v, scatter_dim=2, gather_dim=1)
|
||||
|
||||
if mask_strategy[0] is not None:
|
||||
q = self.tile(q, nccl_info.sp_size).transpose(1, 2).contiguous()
|
||||
k = self.tile(k, nccl_info.sp_size).transpose(1, 2).contiguous()
|
||||
v = self.tile(v, nccl_info.sp_size).transpose(1, 2).contiguous()
|
||||
|
||||
head_num = q.size(1) # 48 // sp_size
|
||||
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)]
|
||||
|
||||
x = sliding_tile_attention(q, k, v, windows, 0, False).transpose(1, 2).contiguous()
|
||||
x = self.untile(x, nccl_info.sp_size)
|
||||
else:
|
||||
x = flash_attn_func(q, k, v, dropout_p=0.0, softmax_scale=None, causal=False)
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
x = all_to_all_4D(x, scatter_dim=1, gather_dim=2)
|
||||
|
||||
x = x.to(q.dtype)
|
||||
return x
|
||||
Executable
+296
@@ -0,0 +1,296 @@
|
||||
# Copyright 2025 StepFun Inc. All Rights Reserved.
|
||||
#
|
||||
# Permission is hereby granted, free of charge, to any person obtaining a copy
|
||||
# of this software and associated documentation files (the "Software"), to deal
|
||||
# in the Software without restriction, including without limitation the rights
|
||||
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
|
||||
# copies of the Software, and to permit persons to whom the Software is
|
||||
# furnished to do so, subject to the following conditions:
|
||||
#
|
||||
# The above copyright notice and this permission notice shall be included in all
|
||||
# copies or substantial portions of the Software.
|
||||
# ==============================================================================
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.models.stepvideo.modules.attentions import Attention
|
||||
from fastvideo.models.stepvideo.modules.normalization import RMSNorm
|
||||
from fastvideo.models.stepvideo.modules.rope import RoPE3D
|
||||
|
||||
|
||||
class SelfAttention(Attention):
|
||||
|
||||
def __init__(self, hidden_dim, head_dim, bias=False, with_rope=True, with_qk_norm=True, attn_type='torch'):
|
||||
super().__init__()
|
||||
self.head_dim = head_dim
|
||||
self.n_heads = hidden_dim // head_dim
|
||||
|
||||
self.wqkv = nn.Linear(hidden_dim, hidden_dim * 3, bias=bias)
|
||||
self.wo = nn.Linear(hidden_dim, hidden_dim, bias=bias)
|
||||
|
||||
self.with_rope = with_rope
|
||||
self.with_qk_norm = with_qk_norm
|
||||
if self.with_qk_norm:
|
||||
self.q_norm = RMSNorm(head_dim, elementwise_affine=True)
|
||||
self.k_norm = RMSNorm(head_dim, elementwise_affine=True)
|
||||
|
||||
if self.with_rope:
|
||||
self.rope_3d = RoPE3D(freq=1e4, F0=1.0, scaling_factor=1.0)
|
||||
self.rope_ch_split = [64, 32, 32]
|
||||
|
||||
self.core_attention = self.attn_processor(attn_type=attn_type)
|
||||
self.parallel = attn_type == 'parallel'
|
||||
|
||||
def apply_rope3d(self, x, fhw_positions, rope_ch_split, parallel=True):
|
||||
x = self.rope_3d(x, fhw_positions, rope_ch_split, parallel)
|
||||
return x
|
||||
|
||||
def forward(self, x, cu_seqlens=None, max_seqlen=None, rope_positions=None, attn_mask=None, mask_strategy=None):
|
||||
xqkv = self.wqkv(x)
|
||||
xqkv = xqkv.view(*x.shape[:-1], self.n_heads, 3 * self.head_dim)
|
||||
|
||||
xq, xk, xv = torch.split(xqkv, [self.head_dim] * 3, dim=-1) ## seq_len, n, dim
|
||||
|
||||
if self.with_qk_norm:
|
||||
xq = self.q_norm(xq)
|
||||
xk = self.k_norm(xk)
|
||||
|
||||
if self.with_rope:
|
||||
xq = self.apply_rope3d(xq, rope_positions, self.rope_ch_split, parallel=self.parallel)
|
||||
xk = self.apply_rope3d(xk, rope_positions, self.rope_ch_split, parallel=self.parallel)
|
||||
|
||||
output = self.core_attention(xq,
|
||||
xk,
|
||||
xv,
|
||||
cu_seqlens=cu_seqlens,
|
||||
max_seqlen=max_seqlen,
|
||||
attn_mask=attn_mask,
|
||||
mask_strategy=mask_strategy)
|
||||
output = rearrange(output, 'b s h d -> b s (h d)')
|
||||
output = self.wo(output)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class CrossAttention(Attention):
|
||||
|
||||
def __init__(self, hidden_dim, head_dim, bias=False, with_qk_norm=True, attn_type='torch'):
|
||||
super().__init__()
|
||||
self.head_dim = head_dim
|
||||
self.n_heads = hidden_dim // head_dim
|
||||
|
||||
self.wq = nn.Linear(hidden_dim, hidden_dim, bias=bias)
|
||||
self.wkv = nn.Linear(hidden_dim, hidden_dim * 2, bias=bias)
|
||||
self.wo = nn.Linear(hidden_dim, hidden_dim, bias=bias)
|
||||
|
||||
self.with_qk_norm = with_qk_norm
|
||||
if self.with_qk_norm:
|
||||
self.q_norm = RMSNorm(head_dim, elementwise_affine=True)
|
||||
self.k_norm = RMSNorm(head_dim, elementwise_affine=True)
|
||||
|
||||
self.core_attention = self.attn_processor(attn_type=attn_type)
|
||||
|
||||
def forward(self, x: torch.Tensor, encoder_hidden_states: torch.Tensor, attn_mask=None):
|
||||
xq = self.wq(x)
|
||||
xq = xq.view(*xq.shape[:-1], self.n_heads, self.head_dim)
|
||||
|
||||
xkv = self.wkv(encoder_hidden_states)
|
||||
xkv = xkv.view(*xkv.shape[:-1], self.n_heads, 2 * self.head_dim)
|
||||
|
||||
xk, xv = torch.split(xkv, [self.head_dim] * 2, dim=-1) ## seq_len, n, dim
|
||||
|
||||
if self.with_qk_norm:
|
||||
xq = self.q_norm(xq)
|
||||
xk = self.k_norm(xk)
|
||||
|
||||
output = self.core_attention(xq, xk, xv, attn_mask=attn_mask)
|
||||
|
||||
output = rearrange(output, 'b s h d -> b s (h d)')
|
||||
output = self.wo(output)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
class GELU(nn.Module):
|
||||
r"""
|
||||
GELU activation function with tanh approximation support with `approximate="tanh"`.
|
||||
|
||||
Parameters:
|
||||
dim_in (`int`): The number of channels in the input.
|
||||
dim_out (`int`): The number of channels in the output.
|
||||
approximate (`str`, *optional*, defaults to `"none"`): If `"tanh"`, use tanh approximation.
|
||||
bias (`bool`, defaults to True): Whether to use a bias in the linear layer.
|
||||
"""
|
||||
|
||||
def __init__(self, dim_in: int, dim_out: int, approximate: str = "none", bias: bool = True):
|
||||
super().__init__()
|
||||
self.proj = nn.Linear(dim_in, dim_out, bias=bias)
|
||||
self.approximate = approximate
|
||||
|
||||
def gelu(self, gate: torch.Tensor) -> torch.Tensor:
|
||||
return torch.nn.functional.gelu(gate, approximate=self.approximate)
|
||||
|
||||
def forward(self, hidden_states):
|
||||
hidden_states = self.proj(hidden_states)
|
||||
hidden_states = self.gelu(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class FeedForward(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
inner_dim: Optional[int] = None,
|
||||
dim_out: Optional[int] = None,
|
||||
mult: int = 4,
|
||||
bias: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
inner_dim = dim * mult if inner_dim is None else inner_dim
|
||||
dim_out = dim if dim_out is None else dim_out
|
||||
self.net = nn.ModuleList([
|
||||
GELU(dim, inner_dim, approximate="tanh", bias=bias),
|
||||
nn.Identity(),
|
||||
nn.Linear(inner_dim, dim_out, bias=bias)
|
||||
])
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor, *args, **kwargs) -> torch.Tensor:
|
||||
for module in self.net:
|
||||
hidden_states = module(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
def modulate(x, scale, shift):
|
||||
x = x * (1 + scale) + shift
|
||||
return x
|
||||
|
||||
|
||||
def gate(x, gate):
|
||||
x = gate * x
|
||||
return x
|
||||
|
||||
|
||||
class StepVideoTransformerBlock(nn.Module):
|
||||
r"""
|
||||
A basic Transformer block.
|
||||
|
||||
Parameters:
|
||||
dim (`int`): The number of channels in the input and output.
|
||||
num_attention_heads (`int`): The number of heads to use for multi-head attention.
|
||||
attention_head_dim (`int`): The number of channels in each head.
|
||||
dropout (`float`, *optional*, defaults to 0.0): The dropout probability to use.
|
||||
cross_attention_dim (`int`, *optional*): The size of the encoder_hidden_states vector for cross attention.
|
||||
activation_fn (`str`, *optional*, defaults to `"geglu"`): Activation function to be used in feed-forward.
|
||||
num_embeds_ada_norm (:
|
||||
obj: `int`, *optional*): The number of diffusion steps used during training. See `Transformer2DModel`.
|
||||
attention_bias (:
|
||||
obj: `bool`, *optional*, defaults to `False`): Configure if the attentions should contain a bias parameter.
|
||||
only_cross_attention (`bool`, *optional*):
|
||||
Whether to use only cross-attention layers. In this case two cross attention layers are used.
|
||||
double_self_attention (`bool`, *optional*):
|
||||
Whether to use two self-attention layers. In this case no cross attention layers are used.
|
||||
upcast_attention (`bool`, *optional*):
|
||||
Whether to upcast the attention computation to float32. This is useful for mixed precision training.
|
||||
norm_elementwise_affine (`bool`, *optional*, defaults to `True`):
|
||||
Whether to use learnable elementwise affine parameters for normalization.
|
||||
norm_type (`str`, *optional*, defaults to `"layer_norm"`):
|
||||
The normalization layer to use. Can be `"layer_norm"`, `"ada_norm"` or `"ada_norm_zero"`.
|
||||
final_dropout (`bool` *optional*, defaults to False):
|
||||
Whether to apply a final dropout after the last feed-forward layer.
|
||||
attention_type (`str`, *optional*, defaults to `"default"`):
|
||||
The type of attention to use. Can be `"default"` or `"gated"` or `"gated-text-image"`.
|
||||
positional_embeddings (`str`, *optional*, defaults to `None`):
|
||||
The type of positional embeddings to apply to.
|
||||
num_positional_embeddings (`int`, *optional*, defaults to `None`):
|
||||
The maximum number of positional embeddings to apply.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
attention_head_dim: int,
|
||||
norm_eps: float = 1e-5,
|
||||
ff_inner_dim: Optional[int] = None,
|
||||
ff_bias: bool = False,
|
||||
attention_type: str = 'parallel'):
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.norm1 = nn.LayerNorm(dim, eps=norm_eps)
|
||||
self.attn1 = SelfAttention(dim,
|
||||
attention_head_dim,
|
||||
bias=False,
|
||||
with_rope=True,
|
||||
with_qk_norm=True,
|
||||
attn_type=attention_type)
|
||||
|
||||
self.norm2 = nn.LayerNorm(dim, eps=norm_eps)
|
||||
self.attn2 = CrossAttention(dim, attention_head_dim, bias=False, with_qk_norm=True, attn_type='torch')
|
||||
|
||||
self.ff = FeedForward(dim=dim, inner_dim=ff_inner_dim, dim_out=dim, bias=ff_bias)
|
||||
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(6, dim) / dim**0.5)
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self,
|
||||
q: torch.Tensor,
|
||||
kv: Optional[torch.Tensor] = None,
|
||||
timestep: Optional[torch.LongTensor] = None,
|
||||
attn_mask=None,
|
||||
rope_positions: list = None,
|
||||
mask_strategy=None) -> torch.Tensor:
|
||||
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = (torch.clone(chunk) for chunk in (
|
||||
self.scale_shift_table[None] + timestep.reshape(-1, 6, self.dim)).chunk(6, dim=1))
|
||||
|
||||
scale_shift_q = modulate(self.norm1(q), scale_msa, shift_msa)
|
||||
|
||||
attn_q = self.attn1(scale_shift_q, rope_positions=rope_positions, mask_strategy=mask_strategy)
|
||||
|
||||
q = gate(attn_q, gate_msa) + q
|
||||
|
||||
attn_q = self.attn2(q, kv, attn_mask)
|
||||
|
||||
q = attn_q + q
|
||||
|
||||
scale_shift_q = modulate(self.norm2(q), scale_mlp, shift_mlp)
|
||||
|
||||
ff_output = self.ff(scale_shift_q)
|
||||
|
||||
q = gate(ff_output, gate_mlp) + q
|
||||
|
||||
return q
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
"""2D Image to Patch Embedding"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
patch_size=64,
|
||||
in_channels=3,
|
||||
embed_dim=768,
|
||||
layer_norm=False,
|
||||
flatten=True,
|
||||
bias=True,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.flatten = flatten
|
||||
self.layer_norm = layer_norm
|
||||
|
||||
self.proj = nn.Conv2d(in_channels,
|
||||
embed_dim,
|
||||
kernel_size=(patch_size, patch_size),
|
||||
stride=patch_size,
|
||||
bias=bias)
|
||||
|
||||
def forward(self, latent):
|
||||
latent = self.proj(latent).to(latent.dtype)
|
||||
if self.flatten:
|
||||
latent = latent.flatten(2).transpose(1, 2) # BCHW -> BNC
|
||||
if self.layer_norm:
|
||||
latent = self.norm(latent)
|
||||
|
||||
return latent
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user