Compare commits

...
67 Commits
Author SHA1 Message Date
JerryZhou54 38ee9dc3b4 Complete Model Config design for VAEs 2025-04-23 19:08:11 +00:00
JerryZhou54 77b013fb8a Add model config for WanVAE 2025-04-23 19:01:13 +00:00
JerryZhou54 c31efe1234 Add model config for VAE 2025-04-23 18:59:25 +00:00
JerryZhou54 c056b89aea Add preliminary design for model config 2025-04-23 18:57:13 +00:00
Kevin Lin eac79b753f [V1] Worker improvements/cleanup (#361) 2025-04-22 00:32:43 -07:00
William Lin 4d58cf20d0 chore: Release FastVideo 0.0.2 and update python requirements (#360) 2025-04-21 14:10:48 -07:00
Kevin LinandWill Lin 52c93ecc9d [V1] Gradio demo with new API (#357)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-04-19 18:14:24 -07:00
William Lin 42d63166ac [V1] Process aware logging; improve logging msg (#356) 2025-04-19 15:03:29 -07:00
William Lin 6db20345a2 [V1] Worker cleanup; Logging clean up; enables isort again (#355) 2025-04-18 19:26:47 -07:00
William Lin ad27ea596c [sta] release 0.0.4 (#354) 2025-04-18 14:54:40 -07:00
William Lin 9aadb4bf8c [1/n] [v1] Add Worker abstractions for User API (#336) 2025-04-18 14:38:46 -07:00
Kevin Lin bd941df271 [Docs] Fix developer guide images (#353) 2025-04-17 22:32:19 -07:00
Yongqi Chen 8a73876d3b add STA to Wan v1 (#349) 2025-04-17 16:35:19 -07:00
Kevin Lin 1483a1138a [CLI] Fix duplicate --num-gpus (#352) 2025-04-17 13:01:48 -07:00
Wei Zhou 5e243d8292 Default to using original WanVAE's encoding/decoding algorithm (#351) 2025-04-17 13:00:25 -07:00
Kevin Lin b0c66d3200 [CI] Docker image improvements (#350) 2025-04-17 12:27:10 -07:00
Wei ZhouandWill Lin c86da2c736 [core] Pipeline config (#343)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-04-15 15:43:10 -07:00
Kevin Lin bae2a19dcf [CI] Add manual trigger to sta-publish and fastvideo-publish (#346) 2025-04-15 15:39:11 -07:00
Kevin Lin 057686f59d [CI] Free up runner disk for sta-publish (#345) 2025-04-15 15:27:09 -07:00
William Lin 67da56628b [STA] Sta release 0.0.3 (#344) 2025-04-15 13:06:03 -07:00
Zhang PeiyuanandSolitaryThinker 2325adffa2 Add STA to V1 (#312)
Co-authored-by: SolitaryThinker <wlsaidhi@gmail.com>
2025-04-15 02:39:58 -07:00
Kevin Lin 20cf836ef1 [CI] Support custom Docker image (#342) 2025-04-14 20:08:53 -07:00
William Lin 2bf69b6f92 [Docs] Initial examples setup and more docs (#332) 2025-04-14 16:53:57 -07:00
William Lin 13583f5ffb [Model] Remove RMSNorm's forward_native hardcode from Wan (#339) 2025-04-14 16:46:43 -07:00
William LinandJerryZhou54 008ee2099a V1 wan rebased (#335)
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
2025-04-11 15:34:54 -07:00
Kevin Lin 137f61f2fe Port tests to v1 (#333) 2025-04-11 01:40:01 -07:00
William Lin ccb262974e [Docs] Add dev guide and doc building CI (#330) 2025-04-09 13:09:00 -07:00
Kevin Lin 30966e3bc9 [CI] Set allowedCudaVersions (#329) 2025-04-09 10:16:05 -07:00
William Lin 7b4272d6b7 [Docs] Fix doc lint (#325) 2025-04-09 10:14:53 -07:00
William Linandkevin314 15553f7706 [CI] Use pre-commit to run linter (#321)
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
2025-04-08 11:38:47 -07:00
William LinandPorridgeSwim 60eeea50bb [Docs] Initial Docs Build (#322)
Co-authored-by: PorridgeSwim <yz3883@columbia.edu>
2025-04-08 11:38:36 -07:00
Kevin Linandkevin314 927b3a40b9 [CI] Add manual triggers for PR workflow (#320)
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
2025-04-07 14:25:39 -07:00
William Linandkevin314 c64f826ae2 Add torch sdpa backend to ssim test (#316)
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
2025-04-07 13:50:46 -07:00
Zhang Peiyuan 55c1040f0b Fix sdpa (#315) 2025-04-06 12:00:34 -07:00
Kevin Linandkevin314 8a3e7aa761 Add ssim test (#314)
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
2025-04-05 17:35:45 -07:00
You ZhouandWill Lin 4324c1c21d refactor the env setup and install of fastvideo (#309)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-04-04 16:33:03 -07:00
Kevin Linandkevin314 2c342ee37f [CI] Add test workflow improvements (#311)
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
2025-04-02 23:14:05 -07:00
Kevin Linandkevin314 708201f531 Set up text encoder tests to work with pytest and Github Actions (#302)
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
2025-04-01 17:56:21 -07:00
1fee098f10 [do not merge] Rebased refactor (#270)
Signed-off-by: <>
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
Co-authored-by: Zhou, Wei <wzhou322@gatech.edu>
Co-authored-by: Kevin Lin <42618777+kevin314@users.noreply.github.com>
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
Co-authored-by: JerryZhou54 <69577934+JerryZhou54@users.noreply.github.com>
Co-authored-by: Yongqi Chen <144848849+BrianChen1129@users.noreply.github.com>
Co-authored-by: Peiyuan Zhang <m2deng@ucsd.edu>
2025-03-29 17:43:47 -05:00
You Zhou 8a77cf22c9 Establish cicd workflow to build and publish FastVideo and STA Kernel (#227) 2025-03-11 20:27:36 -07:00
Yongqi ChenandPeiyuan Zhang d869d90d12 fix training mask strategy issue (#248)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-03-05 20:00:16 -08:00
Zhang Peiyuan 554ee17de5 [BUG] update cfg bug? (#223) 2025-02-27 16:02:44 -08:00
Yongqi ChenandPeiyuan Zhang 0be4fc62c9 fix train/distill issue (#215)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-02-25 08:11:17 -08:00
Yongqi ChenandPeiyuan Zhang 1e08893546 Added multi-GPU support for Hunyuan STA (#211)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-02-21 14:16:28 -08:00
Zhang Peiyuan 09ab452610 Update STA README.md (#206) 2025-02-20 22:26:26 -08:00
Yongqi ChenandPeiyuan Zhang e768b5ec5b Update readme (#202)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-02-20 13:16:25 -08:00
Zhang Peiyuan 59ec42f40e [FIX] Make STA optinal (#204) 2025-02-20 13:09:50 -08:00
rlsu9 5ae5b247b3 [FIX] fix isort format (#203) 2025-02-20 12:20:20 -08:00
ead6c62be4 [Feat] Add STA for StepVideo (#200)
Co-authored-by: rlsu9 <r3su@ucsd.edu>
Co-authored-by: BrianChen1129 <yongqich@umich.edu>
2025-02-20 11:33:58 -08:00
Yongqi ChenandPeiyuan Zhang 6805eaa06c [bug]: fix ori hunyuan inference issue (#199)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-02-19 14:18:15 -08:00
Zhang Peiyuan c39a15551c Update typo (#198) 2025-02-18 19:34:45 -08:00
Zhang Peiyuan e6dda263b0 Update Cite (#195) 2025-02-18 21:01:46 -05:00
Zhang Peiyuan f9482d113c update env (#194) 2025-02-18 20:45:08 -05:00
rlsu9 a3ec969397 [feat]: fix readme demo and add video to readme (#191) 2025-02-18 17:46:32 -05:00
Yongqi ChenandPeiyuan Zhang 76a12cc8a1 Infer sta tea with torch.compile (#190)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-02-18 11:29:36 -08:00
Yongqi ChenandPeiyuan Zhang ac490399c6 fix kernel issue (#185)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-02-16 21:35:56 -08:00
Yongqi ChenandPeiyuan Zhang 9ea39cee57 Add STA and teacache forward (#184)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
2025-02-15 16:22:01 -08:00
Zhang Peiyuanandrlsu9 52e6e612a2 add sliding tile attn (#182)
Co-authored-by: rlsu9 <r3su@ucsd.edu>
2025-02-15 15:44:34 -08:00
Hangliang Ding 9aebc4ada1 Create config.yml (#152) 2025-01-20 20:11:01 -08:00
Yongqi Chen b53cf7425c Lora README update (#155) 2025-01-18 12:30:53 -08:00
Zhang Peiyuan d9ce056901 [typo] 2025-01-13 20:05:57 -08:00
Brian Chen 218449c54d adding hunyuan hf (support lora finetuning); unified hunyuan hf inference with quantization (#135) 2025-01-13 19:47:42 -08:00
Hangliang Ding 221958bcde Update README.md (#131) 2025-01-08 09:02:40 -08:00
Yuzhou Nieand“Peiyuan Zhang” 4a1f1e35bb add parallel for vae decoding (#134)
Co-authored-by: “Peiyuan Zhang” <a1286225768@gmail.com>
2025-01-07 17:14:21 -08:00
rlsu9 e0e05f97f2 [feat]: Add tests for FastVideo (#127) 2025-01-06 12:27:39 -08:00
Zhang Peiyuan dd75ee8509 [Fix] Save CK, Dataset bug fix (#125) 2024-12-31 22:19:10 -08:00
rlsu9 0aed1868df [feat]: Add format auto fixer to main branch (#124) 2024-12-31 15:23:17 -08:00
323 changed files with 1380595 additions and 3856 deletions
+1
View File
@@ -0,0 +1 @@
blank_issues_enabled: false
+250
View File
@@ -0,0 +1,250 @@
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,
required=True,
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"""
# Ensure image name is lowercase (Docker requirement)
image_name = args.image.lower()
print(f"Using specified image: {image_name}")
docker_start_cmd = [
"bash",
"-c",
"apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
]
print(f"Creating RunPod instance with GPU: {args.gpu_type}...")
payload = {
"name": f"fastvideo-{JOB_ID}-{RUN_ID}",
"containerDiskInGb": args.disk_size,
"volumeInGb": args.volume_size,
"gpuTypeIds": [args.gpu_type],
"gpuCount": args.gpu_count,
"imageName": image_name,
"allowedCudaVersions": ["12.4"],
"dockerStartCmd": docker_start_cmd
}
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 = 50
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(20)
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)
# For custom image, we can use the pre-configured environment
setup_steps = [
"tar -xzf /tmp/repo.tar.gz --no-same-owner -C /workspace/",
f"cd /workspace/{repo_name}",
"source /opt/conda/etc/profile.d/conda.sh",
"conda activate fastvideo-dev",
args.test_command
]
remote_command = " && ".join(setup_steps)
ssh_command = [
"ssh", "-o", "StrictHostKeyChecking=no", "-o",
"UserKnownHostsFile=/dev/null", "-o", "ServerAliveInterval=60", "-o",
"ServerAliveCountMax=10", "-p",
str(ssh_port), f"root@{ssh_ip}", remote_command
]
print(f"Connecting to {ssh_ip}:{ssh_port}...")
try:
process = subprocess.Popen(ssh_command,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
universal_newlines=True,
bufsize=0)
stdout_lines = []
print("Command output:")
for line in iter(process.stdout.readline, ''):
print(line.strip())
stdout_lines.append(line)
process.wait()
return_code = process.returncode
success = return_code == 0
stdout_str = "".join(stdout_lines)
if success:
print("Command executed successfully")
else:
print(f"Command failed with exit code {return_code}")
result = {
"success": success,
"return_code": return_code,
"stdout": stdout_str,
"stderr": ""
}
return result
except Exception as e:
print(f"Error executing SSH command: {str(e)}")
result = {"success": False, "error": str(e), "stdout": "", "stderr": ""}
return result
def terminate_pod(pod_id):
"""Terminate the pod"""
print("Terminating RunPod...")
requests.delete(f"{PODS_API}/{pod_id}", headers=HEADERS)
print(f"Terminated pod {pod_id}")
def main():
pod_id = None
try:
pod_id = create_pod()
wait_for_pod(pod_id)
result = execute_command(pod_id)
if result.get("error") is not None:
print(f"Error executing command: {result['error']}")
sys.exit(1)
if not result.get("success", False):
print(
"Tests failed - check the output above for details on which tests failed"
)
sys.exit(1)
finally:
if pod_id:
terminate_pod(pod_id)
if __name__ == "__main__":
main()
+90
View File
@@ -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()
+78
View File
@@ -0,0 +1,78 @@
name: Build and Push Docker Image
on:
workflow_dispatch: # Only manual triggers
jobs:
build-and-push:
runs-on: ubuntu-latest
permissions:
contents: read
packages: write
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Free up disk space
run: |
# Display initial space
echo "Initial disk space:"
df -h
# Remove large directories directly
sudo rm -rf /usr/share/dotnet
sudo rm -rf /usr/local/lib/android
sudo rm -rf /opt/ghc
sudo rm -rf /usr/local/share/boost
sudo rm -rf /usr/share/swift
sudo rm -rf /usr/local/lib/node_modules
sudo rm -rf /usr/local/share/powershell
sudo rm -rf /usr/share/rust
sudo rm -rf /usr/local/.ghcup
# Remove cached files
sudo rm -rf /var/lib/apt/lists/*
sudo rm -rf /var/cache/apt/archives/*
# Clean Docker
docker system prune -af --volumes
# Display available space after cleanup
echo "Disk space after cleanup:"
df -h
- name: Set up Docker Buildx
uses: docker/setup-buildx-action@v3
- name: Login to GitHub Container Registry
uses: docker/login-action@v3
with:
registry: ghcr.io
username: ${{ github.repository_owner }}
password: ${{ secrets.GITHUB_TOKEN }}
- name: Extract metadata for Docker
id: meta
uses: docker/metadata-action@v5
with:
images: ghcr.io/${{ github.repository }}/fastvideo-dev
tags: |
type=raw,value=latest
type=sha,format=short
- name: Build and push Docker image
id: build-push
uses: docker/build-push-action@v6
with:
context: .
push: true
tags: ${{ steps.meta.outputs.tags }}
labels: ${{ steps.meta.outputs.labels }}
cache-from: type=gha
cache-to: type=gha,mode=max
- name: Success message
run: |
echo "✅ Image successfully built and pushed to ghcr.io/${{ github.repository }}/fastvideo-dev:latest"
echo "To run tests with this image, manually trigger the 'Run Tests' workflow."
+83
View File
@@ -0,0 +1,83 @@
# Sample workflow for building and deploying a Hugo site to GitHub Pages
name: Deploy FastVideo Docs to Pages
on:
# Runs on pushes targeting the default branch
push:
branches:
- main
paths:
- "docs/**/*.md"
- "fastvideo/v1/examples/**/*.py"
pull_request:
branches:
- main
types: [opened, ready_for_review, synchronize, reopened]
paths:
- "docs/**/*.md"
- "fastvideo/v1/examples/**/*.py"
# Allows you to run this workflow manually from the Actions tab
workflow_dispatch:
# Sets permissions of the GITHUB_TOKEN to allow deployment to GitHub Pages
permissions:
contents: read
pages: write
id-token: write
# Allow only one concurrent deployment, skipping runs queued between the run in-progress and latest queued.
# However, do NOT cancel in-progress runs as we want to allow these production deployments to complete.
concurrency:
group: "pages"
cancel-in-progress: false
# Default to bash
defaults:
run:
shell: bash
jobs:
pre-commit:
uses: ./.github/workflows/pre-commit.yml
# Build job
build:
runs-on: ubuntu-latest
needs: pre-commit
steps:
- name: Checkout
uses: actions/checkout@v4
- name: Setup Pages
id: pages
uses: actions/configure-pages@v5
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: Install dependencies
run: |
cd docs
pip install -r requirements-docs.txt
- name: Build docs
run: |
cd docs
make clean
make html
- name: Upload artifact
uses: actions/upload-pages-artifact@v3
with:
path: ./docs/build/html
# Deployment job
deploy:
environment:
name: github-pages
url: ${{ steps.deployment.outputs.page_url }}
if: ${{ github.event_name == 'push' }}
runs-on: ubuntu-latest
needs: build
steps:
- name: Deploy to GitHub Pages
id: deployment
uses: actions/deploy-pages@v4
+71
View File
@@ -0,0 +1,71 @@
name: Publish FastVideo to PyPI on Version Change
on:
push:
branches:
- main
paths:
- 'pyproject.toml' # Trigger when pyproject.toml changes
workflow_dispatch:
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' || github.event_name == 'workflow_dispatch' }}
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
}
]
}
]
}
+16
View File
@@ -0,0 +1,16 @@
{
"problemMatcher": [
{
"owner": "mypy",
"pattern": [
{
"regexp": "^(.+):(\\d+):\\s(error|warning):\\s(.+)$",
"file": 1,
"line": 2,
"severity": 3,
"message": 4
}
]
}
]
}
+295
View File
@@ -0,0 +1,295 @@
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:
custom_image:
description: "Custom image from this repository (default: fastvideo-dev:latest)"
required: false
default: "fastvideo-dev:latest"
type: string
run_encoder_test:
description: "Run encoder-test"
required: false
default: false
type: boolean
run_vae_test:
description: "Run vae-test"
required: false
default: false
type: boolean
run_transformer_test:
description: "Run transformer-test"
required: false
default: false
type: boolean
run_ssim_test:
description: "Run ssim-test"
required: false
default: false
type: boolean
env:
PYTHONUNBUFFERED: "1"
concurrency:
group: pr-test-${{ github.ref }}
cancel-in-progress: true
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 }}
vae-test: ${{ steps.filter.outputs.vae-test }}
transformer-test: ${{ steps.filter.outputs.transformer-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/**'
vae-test:
- 'fastvideo/v1/models/vaes/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/vaes/**'
transformer-test:
- 'fastvideo/v1/models/dits/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/transformers/**'
encoder-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.encoder-test == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_encoder_test == 'true')
runs-on: ubuntu-latest
environment: runpod-runners
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: Set up SSH key
run: |
mkdir -p ~/.ssh
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
chmod 600 ~/.ssh/id_rsa
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
- name: Install dependencies
run: pip install requests
- name: Run tests on RunPod
env:
JOB_ID: "encoder-test"
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
timeout-minutes: 30
run: >-
python .github/scripts/runpod_api.py
--gpu-type "NVIDIA A40"
--gpu-count 1
--volume-size 100
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
env:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
JOB_ID: "encoder-test"
run: python .github/scripts/runpod_cleanup.py
vae-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.vae-test == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_vae_test == 'true')
runs-on: ubuntu-latest
environment: runpod-runners
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: Set up SSH key
run: |
mkdir -p ~/.ssh
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
chmod 600 ~/.ssh/id_rsa
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
- name: Install dependencies
run: pip install requests
- name: Run tests on RunPod
env:
JOB_ID: "vae-test"
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
timeout-minutes: 30
run: >-
python .github/scripts/runpod_api.py
--gpu-type "NVIDIA A40"
--gpu-count 1
--volume-size 100
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
env:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
JOB_ID: "vae-test"
run: python .github/scripts/runpod_cleanup.py
transformer-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.transformer-test == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_transformer_test == 'true')
runs-on: ubuntu-latest
environment: runpod-runners
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: Set up SSH key
run: |
mkdir -p ~/.ssh
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
chmod 600 ~/.ssh/id_rsa
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
- name: Install dependencies
run: pip install requests
- name: Run tests on RunPod
env:
JOB_ID: "transformer-test"
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
timeout-minutes: 30
run: >-
python .github/scripts/runpod_api.py
--gpu-type "NVIDIA L40S"
--gpu-count 1
--volume-size 100
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
- name: Terminate RunPod Instances
if: ${{ always() }}
env:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
JOB_ID: "transformer-test"
run: python .github/scripts/runpod_cleanup.py
ssim-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
runs-on: ubuntu-latest
environment: runpod-runners
steps:
- name: Checkout code
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
- name: Set up SSH key
run: |
mkdir -p ~/.ssh
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
chmod 600 ~/.ssh/id_rsa
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
- name: Install dependencies
run: pip install requests
- name: Run tests on RunPod
env:
JOB_ID: "ssim-test"
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
timeout-minutes: 45
run: >-
python .github/scripts/runpod_api.py
--gpu-type "NVIDIA A40"
--gpu-count 2
--disk-size 200
--volume-size 200
--image "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
--test-command "pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
- name: Terminate RunPod Instances
if: ${{ always() }}
env:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
JOB_ID: "ssim-test"
run: python .github/scripts/runpod_cleanup.py
runpod-cleanup:
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
runs-on: ubuntu-latest
steps:
- 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", "vae-test", "transformer-test", "ssim-test"]' # JSON array of job IDs
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
run: python .github/scripts/runpod_cleanup.py
+18
View File
@@ -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
+245
View File
@@ -0,0 +1,245 @@
name: Publish Sliding Tile Attention Kernel to PyPI on Version Change
on:
push:
branches:
- main
paths:
- "csrc/sliding_tile_attention/setup.py"
workflow_dispatch:
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' || github.event_name == 'workflow_dispatch' }}
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: Free up disk space
run: |
echo "Initial disk space:"
df -h
# Remove large directories
sudo rm -rf /usr/share/dotnet
sudo rm -rf /usr/local/lib/android
sudo rm -rf /opt/ghc
sudo rm -rf /usr/local/share/boost
sudo rm -rf /usr/share/swift
sudo rm -rf /usr/local/lib/node_modules
sudo rm -rf /usr/local/share/powershell
sudo rm -rf /usr/share/rust
sudo rm -rf /usr/local/.ghcup
# Remove cached files
sudo rm -rf /var/lib/apt/lists/*
sudo rm -rf /var/cache/apt/archives/*
echo "Disk space after cleanup:"
df -h
- name: Checkout
uses: actions/checkout@v4
- 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' || github.event_name == 'workflow_dispatch' }}
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/
+31
View File
@@ -0,0 +1,31 @@
name: Run Tests
on:
push:
branches: [ main ]
pull_request:
branches: [ main ]
jobs:
test:
runs-on: ubuntu-latest
steps:
- name: Check out repository
uses: actions/checkout@v4
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: '3.12' # or any version you need
- name: Install dependencies
run: |
python -m pip install --upgrade pip setuptools wheel
pip install torch
pip install packaging ninja
pip install -e .
pip install pytest
- name: Run Pytest
run: |
pytest --ignore csrc/sliding_tile_attention/test
+36 -29
View File
@@ -1,17 +1,11 @@
ucf101_stride4x4x4
__pycache__
*.mp4
.ipynb_checkpoints
*.pth
UCF-101/
results/
build/
fastvideo.egg-info/
wandb/
.idea
*.ipynb
*.jpg
*.mp3
*.safetensors
*.mp4
*.png
@@ -20,29 +14,6 @@ wandb/
*.pt
cache_dir/
wandb/
test*
sample_video*
sample_image*
512*
720*
1024*
debug*
private*
caption*
*deepspeed*
revised*
129f*
all*
read*
YSH*
*pick*
*ysh*
hw*
257f*
513f*
taming*
221hw*
65x512x512
runs/
samples/
*validation/
@@ -52,3 +23,39 @@ outputs_video
sbatch.sh
*.out
env
*.o
**/build/
**.pyc
**.txt
**.json
# Distribution / packaging
build/
dist/
*.egg-info/
*.egg
eggs/
.eggs/
# Sphinx documentation
docs/_build/
docs/source/getting_started/examples/
# VSCode
.vscode/
# DS Store
.DS_Store
# vim swap files
*.swo
*.swp
# Python pickle files
*.pkl
# Reference videos
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
# Static images
!docs/source/_static/images/**/*.png
+3
View File
@@ -0,0 +1,3 @@
[submodule "csrc/sliding_tile_attention/tk"]
path = csrc/sliding_tile_attention/tk
url = https://github.com/HazyResearch/ThunderKittens.git
+81
View File
@@ -0,0 +1,81 @@
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/.*|
fastvideo/v1/examples/.*|
.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: 6.0.1
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/" | 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
+48
View File
@@ -0,0 +1,48 @@
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
ENV DEBIAN_FRONTEND=noninteractive
WORKDIR /FastVideo
RUN apt-get update && apt-get install -y --no-install-recommends \
wget \
git \
ca-certificates \
openssh-server \
&& rm -rf /var/lib/apt/lists/*
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
rm Miniconda3-latest-Linux-x86_64.sh
ENV PATH=/opt/conda/bin:$PATH
RUN conda create --name fastvideo-dev python=3.10.0 -y
SHELL ["/bin/bash", "-c"]
# Copy just the pyproject.toml first to leverage Docker cache
COPY pyproject.toml ./
# Create a dummy README to satisfy the installation
RUN echo "# Placeholder" > README.md
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.0.post2 --no-build-isolation && \
conda clean -afy
COPY . .
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
# Remove authentication headers
RUN git config --unset-all http.https://github.com/.extraheader || true
# Set up automatic conda environment activation for all shells
RUN echo 'source /opt/conda/etc/profile.d/conda.sh' >> /root/.bashrc && \
echo 'conda activate fastvideo-dev' >> /root/.bashrc && \
# Ensure .bashrc is sourced for SSH login shells
echo 'if [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
EXPOSE 22
+1 -15
View File
@@ -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.
+150 -42
View File
@@ -4,19 +4,15 @@
FastVideo is a lightweight framework for accelerating large video diffusion models.
https://github.com/user-attachments/assets/5fbc4596-56d6-43aa-98e0-da472cf8e26c
<p align="center">
🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank">FastMochi</a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank">FastHunyuan</a> | 🎮 <a href="https://discord.gg/REBzDQTWWt" target="_blank"> Discord </a> | 🕹️ <a href="https://replicate.com/lucataco/fast-hunyuan-video" target="_blank"> Replicate </a>
</p>
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> <b>Slack</b> </a> |
</p>
https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1
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.
@@ -25,61 +21,95 @@ FastVideo currently offers: (with more to come)
Dev in progress and highly experimental.
## 🎥 More Demos
Fast-Hunyuan comparison with original Hunyuan, achieving an 8X diffusion speed boost with the FastVideo framework.
https://github.com/user-attachments/assets/064ac1d2-11ed-4a0c-955b-4d412a96ef30
Comparison between OpenAI Sora, original Hunyuan and FastHunyuan
https://github.com/user-attachments/assets/d323b712-3f68-42b2-952b-94f6a49c4836
Comparison between original FastHunyuan, LLM-INT8 quantized FastHunyuan and NF4 quantized FastHunyuan
https://github.com/user-attachments/assets/cf89efb5-5f68-4949-a085-f41c1ef26c94
## Change Log
- ```2025/02/20```: FastVideo now supports STA on [StepVideo](https://github.com/stepfun-ai/Step-Video-T2V) with 3.4X speedup!
- ```2025/02/18```: Release the inference code and kernel for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
- ```2025/01/13```: Support Lora finetuning for HunyuanVideo.
- ```2024/12/25```: Enable single 4090 inference for `FastHunyuan`, please rerun the installation steps to update the environment.
- ```2024/12/17```: `FastVideo` v1.0 is released.
## 🔧 Installation from source
The code is tested on Python 3.10-3.12, 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.
You can also install the Sliding Tile Attention package using
```
pip install st_attn==0.0.4
```
## 🚀 Inference
### Inference StepVideo with Sliding Tile Attention
First, download the model:
```
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_diffusers_hunyuan.sh
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
@@ -91,53 +121,100 @@ 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=genmo/mochi-1-preview --local_dir=data/mochi --repo_type=model # original mochi
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model # original hunyuan
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_mochi.sh # for mochi
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 specificed in [Distill Section](#-distill):
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**
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
### ⚡ Lora Finetune
Currently, we only provide Lora Finetune for Mochi model, the command for Lora Finetune is
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
```
bash scripts/finetune/finetune_mochi_lora.sh
```
### Minimum Hardware Requirement
- 40 GB GPU memory each for 2 GPUs with lora
#### 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.
### Finetune with Both Image and Video
Our codebase support finetuning with both image and video.
Currently, both Mochi and Hunyuan models support Lora finetuning through diffusers. To generate personalized videos from your own dataset, you'll need to follow three main steps: dataset preparation, finetuning, and inference.
#### Dataset Preparation
We provide scripts to better help you get started to train on your own characters!
You can run this to organize your dataset to get the videos2caption.json before preprocess. Specify your video folder and corresponding caption folder (caption files should be .txt files and have the same name with its video):
```
python scripts/dataset_preparation/prepare_json_file.py --video_dir data/input_videos/ --prompt_dir data/captions/ --output_path data/output_folder/videos2caption.json --verbose
```
Also, we provide script to resize your videos:
```
python scripts/data_preprocess/resize_videos.py
```
#### Finetuning
After basic dataset preparation and preprocess, you can start to finetune your model using Lora:
```
bash scripts/finetune/finetune_hunyuan_hf_lora.sh
```
#### Inference
For inference with Lora checkpoint, you can run the following scripts with additional parameter `--lora_checkpoint_dir`:
```
bash scripts/inference/inference_hunyuan_hf.sh
```
**We also provide scripts for Mochi in the same directory.**
#### Finetune with Both Image and Video
Our codebase support finetuning with both image and video.
```bash
bash scripts/finetune/finetune_hunyuan.sh
bash scripts/finetune/finetune_mochi_lora_mix.sh
```
For Image-Video Mixture Fine-tuning, make sure to enable the --group_frame option in your script.
For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` option in your script.
## 📑 Development Plan
@@ -149,7 +226,38 @@ For Image-Video Mixture Fine-tuning, make sure to enable the --group_frame optio
- [ ] fp8 support
- [ ] faster load model and save model support
## 🤝 Contributing
We welcome all contributions. Please run `bash format.sh --all` before submitting a pull request.
## 🔧 Testing
Run `pytest` to verify the data preprocessing, checkpoint saving, and sequence parallel pipelines. We recommend adding corresponding test cases in the `test` folder to support your contribution.
## Acknowledgement
We learned and reused code from the following projects: [PCM](https://github.com/G-U-N/Phased-Consistency-Model), [diffusers](https://github.com/huggingface/diffusers), [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan), and [xDiT](https://github.com/xdit-project/xDiT).
We thank MBZUAI and Anyscale for their support throughout this project.
## Citation
If you use FastVideo for your research, please cite our paper:
```bibtex
@misc{zhang2025fastvideogenerationsliding,
title={Fast Video Generation with Sliding Tile Attention},
author={Peiyuan Zhang and Yongqi Chen and Runlong Su and Hangliang Ding and Ion Stoica and Zhenghong Liu and Hao Zhang},
year={2025},
eprint={2502.04507},
archivePrefix={arXiv},
primaryClass={cs.CV},
url={https://arxiv.org/abs/2502.04507},
}
@misc{ding2025efficientvditefficientvideodiffusion,
title={Efficient-vDiT: Efficient Video Diffusion Transformers With Attention Tile},
author={Hangliang Ding and Dacheng Li and Runlong Su and Peiyuan Zhang and Zhijie Deng and Ion Stoica and Hao Zhang},
year={2025},
eprint={2502.06155},
archivePrefix={arXiv},
primaryClass={cs.CV},
url={https://arxiv.org/abs/2502.06155},
}
```
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
File diff suppressed because it is too large Load Diff
Binary file not shown.

After

Width:  |  Height:  |  Size: 751 KiB

+2
View File
@@ -0,0 +1,2 @@
recursive-include tk *
include config.py
+68
View File
@@ -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.
+15
View File
@@ -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'
+76
View File
@@ -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.4"
AUTHOR = "Hao AI Lab"
DESCRIPTION = "Sliding Tile Atteniton Kernel Used in FastVideo"
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/sliding_tile_attention"
# 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"])
+24
View File
@@ -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, int kernel_aspect_ratio_flag
);
#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,46 @@
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, img_latent_shape='30*48*80'):
seq_length = q_all.shape[2]
img_latent_shape_mapping = {
'30x48x80':1,
'36x48x48':2,
'18x48x80':3,
}
if has_text:
assert q_all.shape[
2] >= 115200, "STA currently only supports video with latent size (30, 48, 80), which is 117 frames x 768 x 1280 pixels"
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:
if img_latent_shape == '36x48x48': # Stepvideo 204x768x68
assert q_all.shape[2] == 82944
elif img_latent_shape == '18x48x80': # Wan 69x768x1280
assert q_all.shape[2] == 69120
else:
raise ValueError(f"Unsupported {img_latent_shape}, current shape is {q_all.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
kernel_aspect_ratio_flag = img_latent_shape_mapping[img_latent_shape]
hidden_states = torch.empty_like(q_all)
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
for head_index, (t_kernel, h_kernel, w_kernel) in enumerate(window_size):
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, kernel_aspect_ratio_flag)
if has_text:
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True, kernel_aspect_ratio_flag)
return hidden_states[:, :, :seq_length]
@@ -0,0 +1,831 @@
// # 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, int kernel_aspect_ratio_flag)
{
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_aspect_ratio_flag == 2){
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 1, 1, 6, 6, 6><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} 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;
}
}
else if (kernel_aspect_ratio_flag == 3) {
if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 3) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 1, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 1, 1, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 5) {
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 0, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 0, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 5 && kernel_w_size == 9){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 2, 4, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 2, 4, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 6 && kernel_w_size == 3){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 3, 1, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 3, 1, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 1){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 0, 0, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 1, 0, 0, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 3, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 7){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 3, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 3, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 5 && kernel_w_size == 9){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 2, 4, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false, 0, 2, 4, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 1 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 0, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 0, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 3 && kernel_h_size == 3 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 1, 1, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,1, 1, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 3 && kernel_w_size == 10){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 1, 5, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,0, 1, 5, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else if (kernel_t_size == 1 && kernel_h_size == 6 && kernel_w_size == 5){
cudaFuncSetAttribute(
fwd_attend_ker<128, false, false, false, 0, 3, 2, 3, 6, 10>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
mem_size
);
fwd_attend_ker<128, false, false, false,0, 3, 2, 3, 6, 10><<<grid_image, (32*NUM_WORKERS), mem_size, stream>>>(g);
} else {
// print error
std::cout << "Invalid kernel size" << std::endl;
//print kernel size
std::cout << "Kernel size: " << kernel_t_size << " " << kernel_h_size << " " << kernel_w_size << std::endl;
}
}
else {
std::cout << "Unsupported kernel_aspect_ratio_flag: " << kernel_aspect_ratio_flag << std::endl;
}
}
CHECK_CUDA_ERROR(cudaGetLastError());
cudaStreamSynchronize(stream);
}
return o;
cudaDeviceSynchronize();
}
+151
View File
@@ -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.")
+14 -23
View File
@@ -1,13 +1,15 @@
import argparse
import os
import tempfile
import gradio as gr
import torch
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from diffusers import FlowMatchEulerDiscreteScheduler
from diffusers.utils import export_to_video
from fastvideo.distill.solver import PCMFMScheduler
import tempfile
import os
import argparse
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
def init_args():
@@ -32,7 +34,6 @@ def init_args():
def load_model(args):
device = "cuda" if torch.cuda.is_available() else "cpu"
if args.scheduler_type == "euler":
scheduler = FlowMatchEulerDiscreteScheduler()
else:
@@ -49,13 +50,9 @@ def load_model(args):
if args.transformer_path:
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
else:
transformer = MochiTransformer3DModel.from_pretrained(
args.model_path, subfolder="transformer/"
)
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder="transformer/")
pipe = MochiPipeline.from_pretrained(
args.model_path, transformer=transformer, scheduler=scheduler
)
pipe = MochiPipeline.from_pretrained(args.model_path, transformer=transformer, scheduler=scheduler)
pipe.enable_vae_tiling()
# pipe.to(device)
# if args.cpu_offload:
@@ -76,7 +73,7 @@ def generate_video(
randomize_seed=False,
):
if randomize_seed:
seed = torch.randint(0, 1000000, (1,)).item()
seed = torch.randint(0, 1000000, (1, )).item()
generator = torch.Generator(device="cuda").manual_seed(seed)
@@ -134,9 +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(
@@ -159,9 +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,
@@ -169,9 +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")
@@ -201,4 +192,4 @@ with gr.Blocks() as demo:
)
if __name__ == "__main__":
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
+15
View File
@@ -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
+24
View File
@@ -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"
+20
View File
@@ -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.
+9 -4
View File
@@ -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.
+35
View File
@@ -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
+25
View File
@@ -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
+51
View File
@@ -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.
+8
View File
@@ -0,0 +1,8 @@
.vertical-table-header th.head:not(.stub) {
writing-mode: sideways-lr;
white-space: nowrap;
max-width: 0;
p {
margin: 0;
}
}
+18
View File
@@ -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));
});
});
Binary file not shown.

After

Width:  |  Height:  |  Size: 18 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 27 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 40 KiB

@@ -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>
+260
View File
@@ -0,0 +1,260 @@
# SPDX-License-Identifier: Apache-2.0
# Configuration file for the Sphinx documentation builder.
#
# This file only contains a selection of the most common options. For a full
# list see the documentation:
# https://www.sphinx-doc.org/en/master/usage/configuration.html
# -- Path setup --------------------------------------------------------------
# If extensions (or modules to document with autodoc) are in another directory,
# add these directories to sys.path here. If the directory is relative to the
# documentation root, use os.path.abspath to make it absolute, like shown here.
import datetime
import inspect
import logging
import os
import sys
from typing import Optional
import requests
from sphinx.ext import autodoc
logger = logging.getLogger(__name__)
sys.path.append(os.path.abspath("../.."))
# -- Project information -----------------------------------------------------
project = 'FastVideo'
copyright = f'{datetime.datetime.now().year}, FastVideo Team'
author = 'the FastVideo Team'
# -- General configuration ---------------------------------------------------
# Add any Sphinx extension module names here, as strings. They can be
# extensions coming with Sphinx (named 'sphinx.ext.*') or your custom
# ones.
extensions = [
"sphinx.ext.napoleon",
"sphinx.ext.linkcode",
"sphinx.ext.intersphinx",
"sphinx_copybutton",
"sphinx.ext.autodoc",
"sphinx.ext.autosummary",
"myst_parser",
"sphinxarg.ext",
"sphinx_design",
"sphinx_togglebutton",
]
myst_enable_extensions = [
"colon_fence",
]
# Add any paths that contain templates here, relative to this directory.
templates_path = ['_templates']
# List of patterns, relative to source directory, that match files and
# directories to ignore when looking for source files.
# This pattern also affects html_static_path and html_extra_path.
exclude_patterns: list[str] = ["**/*.template.md", "**/*.inc.md"]
# Exclude the prompt "$" when copying code
copybutton_prompt_text = r"\$ "
copybutton_prompt_is_regexp = True
# -- Options for HTML output -------------------------------------------------
# The theme to use for HTML and HTML Help pages. See the documentation for
# a list of builtin themes.
#
html_title = project
html_theme = 'sphinx_book_theme'
html_logo = '../../assets/logo.jpg'
#html_favicon = 'assets/logos/vllm-logo-only-light.ico'
html_theme_options = {
'path_to_docs': 'docs/source',
'repository_url': 'https://github.com/hao-ai-lab/FastVideo/',
'use_repository_button': True,
'use_edit_page_button': True,
}
# Add any paths that contain custom static files (such as style sheets) here,
# relative to this directory. They are copied after the builtin static files,
# so a file named "default.css" will overwrite the builtin "default.css".
html_static_path = ["_static"]
html_js_files = ["custom.js"]
html_css_files = ["custom.css"]
myst_url_schemes = {
'http': None,
'https': None,
'mailto': None,
'ftp': None,
"gh-issue": {
"url":
"https://github.com/hao-ai-lab/FastVideo/issues/{{path}}#{{fragment}}",
"title": "Issue #{{path}}",
"classes": ["github"],
},
"gh-pr": {
"url":
"https://github.com/hao-ai-lab/FastVideo/pull/{{path}}#{{fragment}}",
"title": "Pull Request #{{path}}",
"classes": ["github"],
},
"gh-dir": {
"url": "https://github.com/hao-ai-lab/FastVideo/tree/main/{{path}}",
"title": "{{path}}",
"classes": ["github"],
},
"gh-file": {
"url": "https://github.com/hao-ai-lab/FastVideo/blob/main/{{path}}",
"title": "{{path}}",
"classes": ["github"],
},
}
# see https://docs.readthedocs.io/en/stable/reference/environment-variables.html # noqa
READTHEDOCS_VERSION_TYPE = os.environ.get('READTHEDOCS_VERSION_TYPE')
if READTHEDOCS_VERSION_TYPE == "tag":
# remove the warning banner if the version is a tagged release
header_file = os.path.join(os.path.dirname(__file__),
"_templates/sections/header.html")
# The file might be removed already if the build is triggered multiple times
# (readthedocs build both HTML and PDF versions separately)
if os.path.exists(header_file):
os.remove(header_file)
# Generate additional rst documentation here.
def setup(app):
from docs.source.generate_examples import generate_examples
generate_examples()
_cached_base: str = ""
_cached_branch: str = ""
def get_repo_base_and_branch(
pr_number: str) -> tuple[Optional[str], Optional[str]]:
global _cached_base, _cached_branch
if _cached_base and _cached_branch:
return _cached_base, _cached_branch
url = f"https://api.github.com/repos/hao-ai-lab/FastVideo/pulls/{pr_number}"
response = requests.get(url)
if response.status_code == 200:
data = response.json()
_cached_base = data['head']['repo']['full_name']
_cached_branch = data['head']['ref']
return _cached_base, _cached_branch
else:
logger.error("Failed to fetch PR details: %s", response)
return None, None
def linkcode_resolve(domain, info):
if domain != 'py':
return None
if not info['module']:
return None
module = info['module']
# try to determine the correct file and line number to link to
obj = sys.modules[module]
# get as specific as we can
lineno: int = 0
filename: str = ""
try:
for part in info['fullname'].split('.'):
obj = getattr(obj, part)
if not (inspect.isclass(obj) or inspect.isfunction(obj)
or inspect.ismethod(obj)):
obj = obj.__class__ # type: ignore[assignment]
lineno = inspect.getsourcelines(obj)[1]
filename = (inspect.getsourcefile(obj)
or f"{filename}.py").split("FastVideo/", 1)[1]
except Exception:
# For some things, like a class member, won't work, so
# we'll use the line number of the parent (the class)
pass
if filename.startswith("checkouts/"):
# a PR build on readthedocs
pr_number = filename.split("/")[1]
filename = filename.split("/", 2)[2]
base, branch = get_repo_base_and_branch(pr_number)
if base and branch:
return f"https://github.com/{base}/blob/{branch}/{filename}#L{lineno}"
# Otherwise, link to the source file on the main branch
return f"https://github.com/hao-ai-lab/FastVideo/blob/main/{filename}#L{lineno}"
# Mock out external dependencies here, otherwise the autodoc pages may be blank.
autodoc_mock_imports = [
"blake3",
"compressed_tensors",
"cpuinfo",
"cv2",
"torch",
"transformers",
"psutil",
"prometheus_client",
"sentencepiece",
"vllm._C",
"PIL",
"numpy",
'triton',
"tqdm",
"tensorizer",
"pynvml",
"outlines",
"xgrammar",
"librosa",
"soundfile",
"gguf",
"lark",
"decord",
]
for mock_target in autodoc_mock_imports:
if mock_target in sys.modules:
logger.info(
"Potentially problematic mock target (%s) found; "
"autodoc_mock_imports cannot mock modules that have already "
"been loaded into sys.modules when the sphinx build starts.",
mock_target)
class MockedClassDocumenter(autodoc.ClassDocumenter):
"""Remove note about base class when a class is derived from object."""
def add_line(self, line: str, source: str, *lineno: int) -> None:
if line == " Bases: :py:class:`object`":
return
super().add_line(line, source, *lineno)
autodoc.ClassDocumenter = MockedClassDocumenter
intersphinx_mapping = {
"python": ("https://docs.python.org/3", None),
"typing_extensions":
("https://typing-extensions.readthedocs.io/en/latest", None),
"aiohttp": ("https://docs.aiohttp.org/en/stable", None),
"pillow": ("https://pillow.readthedocs.io/en/stable", None),
"numpy": ("https://numpy.org/doc/stable", None),
"torch": ("https://pytorch.org/docs/stable", None),
"psutil": ("https://psutil.readthedocs.io/en/stable", None),
}
autodoc_preserve_defaults = True
autodoc_warningiserror = True
navigation_with_keys = False
+138
View File
@@ -0,0 +1,138 @@
(developer-guide)=
# Contributing to FastVideo
Thank you for your interest in contributing to FastVideo. We want to make the process as smooth for you as possible and this is a guide to help get you started!
Our community is open to everyone and welcomes any contributions no matter how large or small.
# Developer Environment:
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only support Linux and CUDA GPUs, but we hope to support other platforms in the future.
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
Install Miniconda:
```
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
bash Miniconda3-latest-Linux-x86_64.sh
source ~/.bashrc
```
Create and activate a Conda environment for FastVideo:
```
conda create -n fastvideo python=3.10 -y
conda activate fastvideo
```
Clone the FastVideo repository and go to the FastVideo directory:
```
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
```
Now you can install FastVideo and setup git hooks for running linting. By using `pre-commit`, the linters will run and have to pass before you'll be able to make a commit.
```bash
pip install -e .[dev]
# Can also install flash-attn (optional)
pip install flash-attn==2.7.0.post2 --no-build-isolation
# Linting, formatting and static type checking
pre-commit install --hook-type pre-commit --hook-type commit-msg
# You can manually run pre-commit with
pre-commit run --all-files
# Unit tests
pytest tests/
```
---
## 🐳 Using the FastVideo Docker Image
If you prefer a containerized development environment or want to avoid managing dependencies manually, you can use our prebuilt Docker image:
**Image:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
### Starting the container
```bash
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
```
This will:
- Start the container with GPU access
- Drop you into a shell with the `fastvideo-dev` Conda environment preconfigured
### Using the container
```bash
# Conda environment should already be active
# FastVideo package installed in editable mode
# Pull the latest changes from remote
cd /FastVideo
git pull
# Run linters and tests
pre-commit run --all-files
pytest tests/
```
---
## 📦 Developing FastVideo on RunPod
You can easily use the FastVideo Docker image as a custom container on [RunPod](https://www.runpod.io) for development or experimentation.
### Creating a new pod
Choose a GPU that supports CUDA 12.4
![RunPod CUDA selection](../_static/images/runpod_cuda.png)
When creating your pod template, use this image:
```
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
```
Paste Container Start Command to support SSH ([RunPod Docs](https://docs.runpod.io/pods/configuration/use-ssh)):
```bash
bash -c "apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
```
![RunPod template configuration](../_static/images/runpod_template.png)
After deploying, the pod will take a few minutes to pull the image and start the SSH service.
![RunPod ssh](../_static/images/runpod_ssh.png)
### Working with the pod
After SSH'ing into your pod, you'll find the `fastvideo-dev` Conda environment already activated.
To pull in the latest changes from the GitHub repo:
```bash
cd /FastVideo
git pull
```
`If you have a persistent volume and want to keep your code changes, you can move /FastVideo to /workspace/FastVideo, or simply clone the repository there.`
Run your development workflows as usual:
```bash
# Run linters
pre-commit run --all-files
# Run tests
pytest tests/
```
+249
View File
@@ -0,0 +1,249 @@
# SPDX-License-Identifier: Apache-2.0
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/docs/source/generate_examples.py
import itertools
import re
from dataclasses import dataclass, field
from pathlib import Path
from typing import Optional
ROOT_DIR = Path(__file__).parent.parent.parent.resolve()
ROOT_DIR_RELATIVE = '../../../..'
EXAMPLE_DIR = ROOT_DIR / "fastvideo/v1/examples"
EXAMPLE_DOC_DIR = ROOT_DIR / "docs/source/getting_started/examples"
def fix_case(text: str) -> str:
subs = {
"api": "API",
"cli": "CLI",
"cpu": "CPU",
"llm": "LLM",
"tpu": "TPU",
"aqlm": "AQLM",
"gguf": "GGUF",
"lora": "LoRA",
"rlhf": "RLHF",
"vllm": "vLLM",
"openai": "OpenAI",
"multilora": "MultiLoRA",
"mlpspeculator": "MLPSpeculator",
r"fp\d+": lambda x: x.group(0).upper(), # e.g. fp16, fp32
r"int\d+": lambda x: x.group(0).upper(), # e.g. int8, int16
}
for pattern, repl in subs.items():
text = re.sub(rf'\b{pattern}\b', repl, text,
flags=re.IGNORECASE) # type: ignore[call-overload]
return text
@dataclass
class Index:
"""
Index class to generate a structured document index.
Attributes:
path (Path): The path save the index file to.
title (str): The title of the index.
description (str): A brief description of the index.
caption (str): An optional caption for the table of contents.
maxdepth (int): The maximum depth of the table of contents. Defaults to 1.
documents (list[str]): A list of document paths to include in the index. Defaults to an empty list.
Methods:
generate() -> str:
Generates the index content as a string in the specified format.
""" # noqa: E501
path: Path
title: str
description: str
caption: str
maxdepth: int = 1
documents: list[str] = field(default_factory=list)
def generate(self) -> str:
content = f"# {self.title}\n\n{self.description}\n\n"
content += ":::{toctree}\n"
content += f":caption: {self.caption}\n:maxdepth: {self.maxdepth}\n"
content += "\n".join(self.documents) + "\n:::\n"
return content
@dataclass
class Example:
"""
Example class for generating documentation content from a given path.
Attributes:
path (Path): The path to the main directory or file.
category (str): The category of the document.
main_file (Path): The main file in the directory.
other_files (list[Path]): list of other files in the directory.
title (str): The title of the document.
Methods:
__post_init__(): Initializes the main_file, other_files, and title attributes.
determine_main_file() -> Path: Determines the main file in the given path.
determine_other_files() -> list[Path]: Determines other files in the directory excluding the main file.
determine_title() -> str: Determines the title of the document.
generate() -> str: Generates the documentation content.
""" # noqa: E501
path: Path
category: Optional[str] = None
main_file: Path = field(init=False)
other_files: list[Path] = field(init=False)
title: str = field(init=False)
def __post_init__(self):
self.main_file = self.determine_main_file()
self.other_files = self.determine_other_files()
self.title = self.determine_title()
def determine_main_file(self) -> Path:
"""
Determines the main file in the given path.
If the path is a file, it returns the path itself. Otherwise, it searches
for Markdown files (*.md) in the directory and returns the first one found.
Returns:
Path: The main file path, either the original path if it's a file or the first
Markdown file found in the directory.
Raises:
IndexError: If no Markdown files are found in the directory.
""" # noqa: E501
return self.path if self.path.is_file() else list(
self.path.glob("*.md")).pop()
def determine_other_files(self) -> list[Path]:
"""
Determine other files in the directory excluding the main file.
This method checks if the given path is a file. If it is, it returns an empty list.
Otherwise, it recursively searches through the directory and returns a list of all
files that are not the main file.
Returns:
list[Path]: A list of Path objects representing the other files in the directory.
""" # noqa: E501
if self.path.is_file():
return []
is_other_file = lambda file: file.is_file() and file != self.main_file
return [file for file in self.path.rglob("*")
if is_other_file(file)] # type: ignore[no-untyped-call]
def determine_title(self) -> str:
return fix_case(self.path.stem.replace("_", " ").title())
def generate(self) -> str:
# Convert the path to a relative path from __file__
make_relative = lambda path: ROOT_DIR_RELATIVE / path.relative_to(
ROOT_DIR)
content = f"Source <gh-file:{self.path.relative_to(ROOT_DIR)}>.\n\n"
include = "include" if self.main_file.suffix == ".md" else \
"literalinclude"
if include == "literalinclude":
content += f"# {self.title}\n\n"
content += f":::{{{include}}} {make_relative(self.main_file)}\n" # type: ignore[no-untyped-call]
if include == "literalinclude":
content += f":language: {self.main_file.suffix[1:]}\n"
content += ":::\n\n"
if not self.other_files:
return content
content += "## Example materials\n\n"
for file in sorted(self.other_files):
include = "include" if file.suffix == ".md" else "literalinclude"
content += f":::{{admonition}} {file.relative_to(self.path)}\n"
content += ":class: dropdown\n\n"
content += f":::{{{include}}} {make_relative(file)}\n:::\n" # type: ignore[no-untyped-call]
content += ":::\n\n"
return content
def generate_examples():
# Create the EXAMPLE_DOC_DIR if it doesn't exist
if not EXAMPLE_DOC_DIR.exists():
EXAMPLE_DOC_DIR.mkdir(parents=True)
# Create empty indices
examples_index = Index(
path=EXAMPLE_DOC_DIR / "examples_index.md",
title="Examples",
description=
"A collection of examples demonstrating usage of FastVideo.\nAll documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.", # noqa: E501
caption="Examples",
maxdepth=2)
# Category indices stored in reverse order because they are inserted into
# examples_index.documents at index 0 in order
category_indices = {
# "other":
# Index(
# path=EXAMPLE_DOC_DIR / "examples_other_index.md",
# title="Other",
# description=
# "Other examples that don't strongly fit into the online or offline serving categories.", # noqa: E501
# caption="Examples",
# ),
# "online_serving":
# Index(
# path=EXAMPLE_DOC_DIR / "examples_online_serving_index.md",
# title="Online Serving",
# description=
# "Online serving examples demonstrate how to use FastVideo in an online setting, where the model is queried for predictions in real-time.", # noqa: E501
# caption="Examples",
# ),
"inference":
Index(
path=EXAMPLE_DOC_DIR / "examples_inference_index.md",
title="Inference",
description=
"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:
print(category)
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):
print(example)
doc_path = EXAMPLE_DOC_DIR / f"{example.path.stem}.md"
with open(doc_path, "w+") as f:
f.write(example.generate())
# Add the example to the appropriate index
assert example.category is not None
index = category_indices.get(example.category, examples_index)
index.documents.append(example.path.stem)
# 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,95 @@
(fastvideo-installation)=
# 🔧 Installation
FastVideo currently only supports Linux and NVIDIA CUDA GPUs.
FastVideo has been tested on the following GPUs, but it should work on any GPUs that supports CUDA 12.4+, please create an issue if you discover any issues:
- RTX 4090
- A40
- L40S
- A100
- H100
## Requirements
- OS: Linux
- Python: 3.10-3.12
- CUDA 12.4+ (Untested on CUDA < 12.4)
## Installation Options
### Option 1: Quick Install
```bash
pip install fastvideo
```
### Option 2: Installation from Source
We recommend using a Python environment such as Conda.
#### 1. [Optional] Install Miniconda (if not already installed)
```bash
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
bash Miniconda3-latest-Linux-x86_64.sh
source ~/.bashrc
```
#### 2. [Optional] Create and activate a Conda environment for FastVideo
```bash
conda create -n fastvideo python=3.10 -y
conda activate fastvideo
```
#### 3. Clone the FastVideo repository
```bash
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
```
#### 4. Install FastVideo
Basic installation:
```bash
pip install -e .
```
## Optional Dependencies
### Flash Attention
```bash
pip install flash-attn==2.7.0.post2 --no-build-isolation
```
### Sliding Tile Attention (STA) (Requires CUDA 12.4+ and H100)
To try Sliding Tile Attention (optional), please follow the instructions in [csrc/sliding_tile_attention/README.md](#sta-installation) to install STA.
## Development Environment Setup
If you're planning to contribute to FastVideo please see the following page:
[Contributor Guide](#developer-guide)
## Hardware Requirements
### For Basic Inference
- NVIDIA GPU with CUDA support
- Minimum 20GB VRAM for quantized models (e.g., single RTX 4090)
### For Lora Finetuning
- 40GB GPU memory each for 2 GPUs with lora
- 30GB GPU memory each for 2 GPUs with CPU offload and lora
### For Full Finetuning/Distillation
- Multiple high-memory GPUs recommended (e.g., H100)
## Troubleshooting
If you encounter any issues during installation, please open an issue on our [GitHub repository](https://github.com/hao-ai-lab/FastVideo).
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg) for additional support.
+89
View File
@@ -0,0 +1,89 @@
# 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/wanvideo
inference/stepvideo
inference/hunyuanvideo
inference/fasthunyuan
inference/fastmochi
:::
:::{toctree}
:caption: Developer Guide
:maxdepth: 1
developer_guide/overview
:::
## Indices and tables
- {ref}`genindex`
- {ref}`modindex`
+33
View File
@@ -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).
+9
View File
@@ -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
+18
View File
@@ -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.
+16
View File
@@ -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
```
+44
View File
@@ -0,0 +1,44 @@
(wanvideo)=
# WanVideo
## Inference T2V with WanVideo
First, download the model:
```bash
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-T2V-1.3B-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
```
or
```bash
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-T2V-14B-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
```
Then run the inference using:
```bash
sh scripts/inference/v1_inference_wan.sh
```
Remember to set `MODEL_BASE` and `num_gpus` accordingly.
## Inference I2V with WanVideo
First, download the model:
```bash
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-I2V-14B-480P-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
```
or
```bash
python scripts/huggingface/download_hf.py --repo_id=Wan-AI/Wan2.1-I2V-14B-720P-Diffusers --local_dir=YOUR_LOCAL_DIR --repo_type=model
```
Then run the inference using:
```bash
sh scripts/inference/v1_inference_wan_i2v.sh
```
Remember to set `MODEL_BASE` and `num_gpus` accordingly.
@@ -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)
```
-10
View File
@@ -1,10 +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
# install fastvideo
pip install -e .
+3
View File
@@ -0,0 +1,3 @@
# Basic
The class provides the main python interface for using FastVideo's inference pipeline.
+1
View File
@@ -0,0 +1 @@
print('Hello, world!')
+3
View File
@@ -0,0 +1,3 @@
from fastvideo.v1.entrypoints.video_generator import VideoGenerator
__all__ = ["VideoGenerator"]
@@ -1,24 +1,27 @@
import argparse
import torch
from accelerate.logging import get_logger
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
from diffusers.utils import export_to_video
import json
import os
import torch
import torch.distributed as dist
from accelerate.logging import get_logger
from diffusers.utils import export_to_video
from diffusers.video_processor import VideoProcessor
from torch.utils.data import DataLoader, Dataset
from torch.utils.data.distributed import DistributedSampler
from tqdm import tqdm
from fastvideo.utils.load import load_text_encoder, load_vae
logger = get_logger(__name__)
from torch.utils.data import Dataset
from torch.utils.data.distributed import DistributedSampler
from torch.utils.data import DataLoader
from fastvideo.utils.load import load_text_encoder, load_vae
from diffusers.video_processor import VideoProcessor
from tqdm import tqdm
class T5dataset(Dataset):
def __init__(
self, json_path, vae_debug,
self,
json_path,
vae_debug,
):
self.json_path = json_path
self.vae_debug = vae_debug
@@ -32,9 +35,7 @@ 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:
@@ -54,9 +55,7 @@ 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)
@@ -70,9 +69,7 @@ def main(args):
text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
vae.enable_tiling()
sampler = DistributedSampler(
train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True
)
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
train_dataloader = DataLoader(
train_dataset,
sampler=sampler,
@@ -84,23 +81,16 @@ def main(args):
for _, data in tqdm(enumerate(train_dataloader), disable=local_rank != 0):
with torch.inference_mode():
with torch.autocast("cuda", dtype=autocast_type):
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(
prompt=data["caption"],
)
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(prompt=data["caption"], )
if args.vae_debug:
latents = data["latents"]
video = vae.decode(latents.to(device), return_dict=False)[0]
video = videoprocessor.postprocess_video(video)
for idx, video_name in enumerate(data["filename"]):
prompt_embed_path = os.path.join(
args.output_dir, "prompt_embed", video_name + ".pt"
)
video_path = os.path.join(
args.output_dir, "video", video_name + ".mp4"
)
prompt_attention_mask_path = os.path.join(
args.output_dir, "prompt_attention_mask", video_name + ".pt"
)
prompt_embed_path = os.path.join(args.output_dir, "prompt_embed", video_name + ".pt")
video_path = os.path.join(args.output_dir, "video", video_name + ".mp4")
prompt_attention_mask_path = os.path.join(args.output_dir, "prompt_attention_mask",
video_name + ".pt")
# save latent
torch.save(prompt_embeds[idx], prompt_embed_path)
torch.save(prompt_attention_mask[idx], prompt_attention_mask_path)
@@ -1,19 +1,17 @@
from fastvideo.dataset import getdataset
from torch.utils.data import DataLoader
from fastvideo.utils.dataset_utils import Collate
import argparse
import torch
from accelerate import Accelerator
from accelerate.logging import get_logger
from accelerate.utils import ProjectConfiguration
import json
import os
from diffusers import AutoencoderKLMochi
import torch
import torch.distributed as dist
from accelerate.logging import get_logger
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from fastvideo.utils.load import load_vae
from tqdm import tqdm
from fastvideo.dataset import getdataset
from fastvideo.utils.load import load_vae
logger = get_logger(__name__)
@@ -22,9 +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,
@@ -32,12 +28,10 @@ def main(args):
num_workers=args.dataloader_num_workers,
)
encoder_device = torch.device(f"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)
@@ -47,14 +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]
@@ -91,9 +81,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("--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)
@@ -119,10 +107,8 @@ if __name__ == "__main__":
"--logging_dir",
type=str,
default="logs",
help=(
"[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."
),
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
args = parser.parse_args()
@@ -1,19 +1,13 @@
import argparse
import torch
from accelerate.logging import get_logger
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
from diffusers.utils import export_to_video
import json
import os
import torch
import torch.distributed as dist
from accelerate.logging import get_logger
from fastvideo.utils.load import load_text_encoder
logger = get_logger(__name__)
from torch.utils.data import Dataset
from torch.utils.data.distributed import DistributedSampler
from torch.utils.data import DataLoader
from fastvideo.utils.load import load_text_encoder, load_vae
from diffusers.video_processor import VideoProcessor
from tqdm import tqdm
def main(args):
@@ -24,9 +18,7 @@ def main(args):
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
torch.cuda.set_device(local_rank)
if not dist.is_initialized():
dist.init_process_group(
backend="nccl", init_method="env://", world_size=world_size, rank=local_rank
)
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
autocast_type = torch.float16 if args.model_type == "hunyuan" else torch.bfloat16
@@ -37,23 +29,17 @@ 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
)
json_data = []
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()
prompts = [line.strip() for line in lines]
for prompt in prompts:
with torch.inference_mode():
with torch.autocast("cuda", dtype=autocast_type):
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(
prompt
)
prompt_embeds, prompt_attention_mask = text_encoder.encode_prompt(prompt)
file_name = prompt.split(".")[0]
prompt_embed_path = os.path.join(
args.output_dir, "validation", "prompt_embed", f"{file_name}.pt"
)
prompt_embed_path = os.path.join(args.output_dir, "validation", "prompt_embed", f"{file_name}.pt")
prompt_attention_mask_path = os.path.join(
args.output_dir,
"validation",
+18 -32
View File
@@ -1,14 +1,9 @@
from transformers import AutoTokenizer
from torchvision import transforms
from torchvision.transforms import Lambda
from transformers import AutoTokenizer
from fastvideo.dataset.t2v_datasets import T2V_dataset
from fastvideo.dataset.latent_datasets import LatentDataset
from fastvideo.dataset.transform import (
Normalize255,
TemporalRandomCrop,
CenterCropResizeVideo,
)
from fastvideo.dataset.transform import CenterCropResizeVideo, Normalize255, TemporalRandomCrop
def getdataset(args):
@@ -20,26 +15,17 @@ def getdataset(args):
resize = [
CenterCropResizeVideo((args.max_height, args.max_width)),
]
transform = transforms.Compose(
[
# Normalize255(),
*resize,
# RandomHorizontalFlipVideo(p=0.5), # in case their caption have position decription
# norm_fun
]
)
transform_topcrop = transforms.Compose(
[
Normalize255(),
*resize_topcrop,
# RandomHorizontalFlipVideo(p=0.5), # in case their caption have position decription
norm_fun,
]
)
transform = transforms.Compose([
# Normalize255(),
*resize,
])
transform_topcrop = transforms.Compose([
Normalize255(),
*resize_topcrop,
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,
@@ -53,11 +39,13 @@ def getdataset(args):
if __name__ == "__main__":
from accelerate import Accelerator
from fastvideo.dataset.t2v_datasets import dataset_prog
import random
from accelerate import Accelerator
from tqdm import tqdm
from fastvideo.dataset.t2v_datasets import dataset_prog
args = type(
"args",
(),
@@ -92,9 +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:
+17 -22
View File
@@ -1,13 +1,18 @@
import torch
from torch.utils.data import Dataset
import json
import os
import random
import torch
from torch.utils.data import Dataset
class LatentDataset(Dataset):
def __init__(
self, json_path, num_latent_t, cfg_rate,
self,
json_path,
num_latent_t,
cfg_rate,
):
# data_merge_path: video_dir, latent_dir, prompt_embed_dir, json_path
self.json_path = json_path
@@ -16,9 +21,7 @@ class LatentDataset(Dataset):
self.video_dir = os.path.join(self.datase_dir_path, "video")
self.latent_dir = os.path.join(self.datase_dir_path, "latent")
self.prompt_embed_dir = os.path.join(self.datase_dir_path, "prompt_embed")
self.prompt_attention_mask_dir = os.path.join(
self.datase_dir_path, "prompt_attention_mask"
)
self.prompt_attention_mask_dir = os.path.join(self.datase_dir_path, "prompt_attention_mask")
with open(self.json_path, "r") as f:
self.data_anno = json.load(f)
# json.load(f) already keeps the order
@@ -28,10 +31,7 @@ class LatentDataset(Dataset):
self.uncond_prompt_embed = torch.zeros(256, 4096).to(torch.float32)
# 256 zeros
self.uncond_prompt_mask = torch.zeros(256).bool()
self.lengths = [
data_item["length"] if "length" in data_item else 1
for data_item in self.data_anno
]
self.lengths = [data_item["length"] if "length" in data_item else 1 for data_item in self.data_anno]
def __getitem__(self, idx):
latent_file = self.data_anno[idx]["latent_path"]
@@ -43,7 +43,7 @@ class LatentDataset(Dataset):
map_location="cpu",
weights_only=True,
)
latent = latent.squeeze(0)[:, -self.num_latent_t :]
latent = latent.squeeze(0)[:, -self.num_latent_t:]
if random.random() < self.cfg_rate:
prompt_embed = self.uncond_prompt_embed
prompt_attention_mask = self.uncond_prompt_mask
@@ -54,9 +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,
)
@@ -89,16 +87,15 @@ def latent_collate_function(batch):
0,
max_w - latent.shape[3],
),
)
for latent in latents
) for latent in latents
]
# attn mask
latent_attn_mask = torch.ones(len(latents), max_t, max_h, max_w)
# set to 0 if padding
for i, latent in enumerate(latents):
latent_attn_mask[i, latent.shape[1] :, :, :] = 0
latent_attn_mask[i, :, latent.shape[2] :, :] = 0
latent_attn_mask[i, :, :, latent.shape[3] :] = 0
latent_attn_mask[i, latent.shape[1]:, :, :] = 0
latent_attn_mask[i, :, latent.shape[2]:, :] = 0
latent_attn_mask[i, :, :, latent.shape[3]:] = 0
prompt_embeds = torch.stack(prompt_embeds, dim=0)
prompt_attention_masks = torch.stack(prompt_attention_masks, dim=0)
@@ -108,9 +105,7 @@ 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
)
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,
+27 -50
View File
@@ -1,18 +1,18 @@
import json
import os, io, csv, math, random
import numpy as np
from einops import rearrange
from decord import VideoReader
from os.path import join as opj
import math
import os
import random
from collections import Counter
from os.path import join as opj
import numpy as np
import torch
from torch.utils.data.dataset import Dataset
from torch.utils.data import DataLoader, Dataset, get_worker_info
from tqdm import tqdm
from PIL import Image
from fastvideo.utils.dataset_utils import DecordInit
import torchvision
from einops import rearrange
from PIL import Image
from torch.utils.data import Dataset
from fastvideo.utils.dataset_utils import DecordInit
from fastvideo.utils.logging_ import main_print
@@ -27,6 +27,7 @@ class SingletonMeta(type):
class DataSetProg(metaclass=SingletonMeta):
def __init__(self):
self.cap_list = []
self.elements = []
@@ -56,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
@@ -73,6 +72,7 @@ def filter_resolution(h, w, max_h_div_w_ratio=17 / 16, min_h_div_w_ratio=8 / 16)
class T2V_dataset(Dataset):
def __init__(self, args, transform, temporal_sample, tokenizer, transform_topcrop):
self.data = args.data_merge_path
self.num_frames = args.num_frames
@@ -92,7 +92,7 @@ class T2V_dataset(Dataset):
self.v_decoder = DecordInit()
self.video_length_tolerance_range = args.video_length_tolerance_range
self.support_Chinese = True
if not ("mt5" in args.text_encoder_name):
if "mt5" not in args.text_encoder_name:
self.support_Chinese = False
cap_list = self.get_cap_list()
@@ -129,9 +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")
@@ -180,20 +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)
) # [1 C H W] -> num_img [1 C H W]
image = (self.transform_topcrop(image) if "human_images" in image_data["path"] else self.transform(image)
) # [1 C H W] -> num_img [1 C H W]
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
image = image.float() / 127.5 - 1.0
caps = (
image_data["cap"]
if isinstance(image_data["cap"], list)
else [image_data["cap"]]
)
caps = (image_data["cap"] if isinstance(image_data["cap"], list) else [image_data["cap"]])
caps = [random.choice(caps)]
text = caps
input_ids, cond_mask = [], []
@@ -247,10 +238,7 @@ 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"]
@@ -271,23 +259,18 @@ 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
@@ -298,9 +281,7 @@ class T2V_dataset(Dataset):
# frame_indices = frame_indices[:self.num_frames] # head crop
i["sample_frame_index"] = frame_indices.tolist()
new_cap_list.append(i)
i["sample_num_frames"] = len(
i["sample_frame_index"]
) # will use in dataloader(group sampler)
i["sample_num_frames"] = len(i["sample_frame_index"]) # will use in dataloader(group sampler)
sample_num_frames.append(i["sample_num_frames"])
elif path.endswith(".jpg"): # image
cnt_img += 1
@@ -309,15 +290,13 @@ class T2V_dataset(Dataset):
sample_num_frames.append(i["sample_num_frames"])
else:
raise NameError(
f"Unknown file extention {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):
@@ -330,9 +309,7 @@ class T2V_dataset(Dataset):
def read_jsons(self, data):
cap_lists = []
with open(data, "r") as f:
folder_anno = [
i.strip().split(",") for i in f.readlines() if len(i.strip()) > 0
]
folder_anno = [i.strip().split(",") for i in f.readlines() if len(i.strip()) > 0]
print(folder_anno)
for folder, anno in folder_anno:
with open(anno, "r") as f:
+59 -84
View File
@@ -1,7 +1,8 @@
import torch
import random
import numbers
from torchvision.transforms import RandomCrop, RandomResizedCrop
import random
import torch
from PIL import Image
def _is_tensor_video_clip(clip):
@@ -20,21 +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):
@@ -44,14 +39,12 @@ def crop(clip, i, j, h, w):
"""
if len(clip.size()) != 4:
raise ValueError("clip should be a 4D tensor")
return clip[..., i : i + h, j : j + w]
return clip[..., i:i + h, j:j + w]
def resize(clip, target_size, interpolation_mode):
if len(target_size) != 2:
raise ValueError(
f"target size should be tuple (height, width), instead got {target_size}"
)
raise ValueError(f"target size should be tuple (height, width), instead got {target_size}")
return torch.nn.functional.interpolate(
clip,
size=target_size,
@@ -63,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(
@@ -153,16 +144,14 @@ def random_shift_crop(clip):
h, w = clip.size(-2), clip.size(-1)
if h <= w:
long_edge = w
short_edge = h
else:
long_edge = h
short_edge = w
th, tw = short_edge, short_edge
i = torch.randint(0, h - th + 1, size=(1,)).item()
j = torch.randint(0, w - tw + 1, size=(1,)).item()
i = torch.randint(0, h - th + 1, size=(1, )).item()
j = torch.randint(0, w - tw + 1, size=(1, )).item()
return crop(clip, i, j, th, tw)
@@ -177,9 +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
@@ -217,6 +204,7 @@ def hflip(clip):
class RandomCropVideo:
def __init__(self, size):
if isinstance(size, numbers.Number):
self.size = (int(size), int(size))
@@ -239,15 +227,13 @@ class RandomCropVideo:
th, tw = self.size
if h < th or w < tw:
raise ValueError(
f"Required crop size {(th, tw)} is larger than input image size {(h, w)}"
)
raise ValueError(f"Required crop size {(th, tw)} is larger than input image size {(h, w)}")
if w == tw and h == th:
return 0, 0, h, w
i = torch.randint(0, h - th + 1, size=(1,)).item()
j = torch.randint(0, w - tw + 1, size=(1,)).item()
i = torch.randint(0, h - th + 1, size=(1, )).item()
j = torch.randint(0, w - tw + 1, size=(1, )).item()
return i, j, th, tw
@@ -256,6 +242,7 @@ class RandomCropVideo:
class SpatialStrideCropVideo:
def __init__(self, stride):
self.stride = stride
@@ -288,7 +275,10 @@ class LongSideResizeVideo:
"""
def __init__(
self, size, skip_low_resolution=False, interpolation_mode="bilinear",
self,
size,
skip_low_resolution=False,
interpolation_mode="bilinear",
):
self.size = size
self.skip_low_resolution = skip_low_resolution
@@ -311,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:
@@ -327,12 +315,13 @@ class CenterCropResizeVideo:
"""
def __init__(
self, size, top_crop=False, interpolation_mode="bilinear",
self,
size,
top_crop=False,
interpolation_mode="bilinear",
):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}"
)
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
self.top_crop = top_crop
self.interpolation_mode = interpolation_mode
@@ -346,9 +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,
@@ -368,13 +355,13 @@ class UCFCenterCropVideo:
"""
def __init__(
self, size, interpolation_mode="bilinear",
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}"
)
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
@@ -389,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
@@ -405,13 +390,13 @@ class KineticsRandomCropResizeVideo:
"""
def __init__(
self, size, interpolation_mode="bilinear",
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}"
)
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
@@ -425,14 +410,15 @@ class KineticsRandomCropResizeVideo:
class CenterCropVideo:
def __init__(
self, size, interpolation_mode="bilinear",
self,
size,
interpolation_mode="bilinear",
):
if isinstance(size, tuple):
if len(size) != 2:
raise ValueError(
f"size should be tuple (height, width), instead got {size}"
)
raise ValueError(f"size should be tuple (height, width), instead got {size}")
self.size = size
else:
self.size = (size, size)
@@ -559,9 +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
@@ -569,27 +553,22 @@ class DynamicSampleDuration(object):
if __name__ == "__main__":
from torchvision import transforms
import torchvision.io as io
import numpy as np
from torchvision.utils import save_image
import os
vframes, aframes, info = io.read_video(
filename="./v_Archery_g01_c03.avi", pts_unit="sec", output_format="TCHW"
)
import numpy as np
import torchvision.io as io
from torchvision import transforms
from torchvision.utils import save_image
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
),
]
)
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),
])
target_video_len = 32
frame_interval = 1
@@ -603,9 +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]
@@ -616,9 +593,7 @@ 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)
+112 -234
View File
@@ -1,53 +1,41 @@
# !/bin/python3
# isort: skip_file
import argparse
import math
import os
from fastvideo.utils.parallel_states import (
initialize_sequence_parallel_state,
destroy_sequence_parallel_group,
get_sequence_parallel_state,
nccl_info,
)
from fastvideo.utils.communications import sp_parallel_dataloader_wrapper, broadcast
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.utils.validation import log_validation
import time
from torch.utils.data import DataLoader
from collections import deque
from copy import deepcopy
import torch
from torch.distributed.fsdp import ShardingStrategy
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
StateDictType,
FullStateDictConfig,
)
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
import json
from torch.utils.data.distributed import DistributedSampler
from fastvideo.utils.dataset_utils import LengthGroupedSampler
import torch.distributed as dist
import wandb
from accelerate.utils import set_seed
from tqdm.auto import tqdm
from fastvideo.utils.fsdp_util import get_dit_fsdp_kwargs, apply_fsdp_checkpointing
from diffusers import FlowMatchEulerDiscreteScheduler
from fastvideo.utils.load import load_transformer
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from copy import deepcopy
from diffusers.optimization import get_scheduler
from diffusers.utils import check_min_version
from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_function
import torch.distributed as dist
from safetensors.torch import save_file
from peft import LoraConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from fastvideo.utils.checkpoint import (
save_checkpoint,
save_lora_checkpoint,
resume_lora_optimizer,
)
from torch.distributed.fsdp import ShardingStrategy
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
from fastvideo.utils.checkpoint import (resume_lora_optimizer, save_checkpoint, save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast, sp_parallel_dataloader_wrapper)
from fastvideo.utils.dataset_utils import LengthGroupedSampler
from fastvideo.utils.fsdp_util import (apply_fsdp_checkpointing, get_dit_fsdp_kwargs)
from fastvideo.utils.load import load_transformer
from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group, get_sequence_parallel_state,
initialize_sequence_parallel_state)
from fastvideo.utils.validation import log_validation
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
check_min_version("0.31.0")
import time
from collections import deque
def main_print(content):
@@ -55,31 +43,6 @@ def main_print(content):
print(content)
def save_checkpoint(transformer, rank, output_dir, step):
main_print(f"--> saving checkpoint at step {step}")
with FSDP.state_dict_type(
transformer,
StateDictType.FULL_STATE_DICT,
FullStateDictConfig(offload_to_cpu=True, rank0_only=True),
):
cpu_state = transformer.state_dict()
# todo move to get_state_dict
if rank <= 0:
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
os.makedirs(save_dir, exist_ok=True)
# save using safetensors
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
save_file(cpu_state, weight_path)
config_dict = dict(transformer.config)
if "dtype" in config_dict:
del config_dict["dtype"] # TODO
config_path = os.path.join(save_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
main_print(f"--> checkpoint saved at step {step}")
def reshard_fsdp(model):
for m in FSDP.fsdp_modules(model):
if m._has_params and m.sharding_strategy is not ShardingStrategy.NO_SHARD:
@@ -88,17 +51,15 @@ def reshard_fsdp(model):
def get_norm(model_pred, norms, gradient_accumulation_steps):
fro_norm = (
torch.linalg.matrix_norm(model_pred, ord="fro") / gradient_accumulation_steps
)
largest_singular_value = (
torch.linalg.matrix_norm(model_pred, ord=2) / gradient_accumulation_steps
)
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
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()
norms["fro"] += torch.mean(fro_norm).item() # codespell:ignore
norms["largest singular value"] += torch.mean(largest_singular_value).item()
norms["absolute mean"] += absolute_mean.item()
norms["absolute max"] += absolute_max.item()
@@ -132,7 +93,7 @@ def distill_one_step(
total_loss = 0.0
optimizer.zero_grad()
model_pred_norm = {
"fro": 0.0,
"fro": 0.0, # codespell:ignore
"largest singular value": 0.0,
"absolute mean": 0.0,
"absolute max": 0.0,
@@ -147,9 +108,7 @@ 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.
@@ -160,9 +119,7 @@ def distill_one_step(
timesteps = (sigmas * noise_scheduler.config.num_train_timesteps).view(-1)
# if squeeze to [], unsqueeze to [1]
timesteps_prev = (
sigmas_prev * noise_scheduler.config.num_train_timesteps
).view(-1)
timesteps_prev = (sigmas_prev * noise_scheduler.config.num_train_timesteps).view(-1)
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
# Predict the noise residual
with torch.autocast("cuda", dtype=torch.bfloat16):
@@ -174,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):
@@ -205,9 +160,7 @@ 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
)
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
@@ -230,42 +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")
# 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()
@@ -275,13 +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()
@@ -312,7 +245,7 @@ def main(args):
if rank <= 0 and args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
# For mixed precision training we cast all non-trainable weigths to half-precision
# For mixed precision training we cast all non-trainable weights to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
# Create model:
@@ -344,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,
@@ -364,23 +294,26 @@ def main(args):
transformer._no_split_modules = no_split_modules
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
transformer = FSDP(transformer, **fsdp_kwargs,)
teacher_transformer = FSDP(teacher_transformer, **fsdp_kwargs,)
transformer = FSDP(
transformer,
**fsdp_kwargs,
)
teacher_transformer = FSDP(
teacher_transformer,
**fsdp_kwargs,
)
if args.use_ema:
ema_transformer = FSDP(ema_transformer, **fsdp_kwargs,)
main_print(f"--> model loaded")
ema_transformer = FSDP(
ema_transformer,
**fsdp_kwargs,
)
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)
@@ -388,9 +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,
@@ -418,9 +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
@@ -437,20 +367,15 @@ def main(args):
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t, args.cfg)
uncond_prompt_embed = train_dataset.uncond_prompt_embed
uncond_prompt_mask = train_dataset.uncond_prompt_mask
sampler = (
LengthGroupedSampler(
args.train_batch_size,
rank=rank,
world_size=world_size,
lengths=train_dataset.lengths,
group_frame=args.group_frame,
group_resolution=args.group_resolution,
)
if (args.group_frame or args.group_resolution)
else DistributedSampler(
train_dataset, rank=rank, num_replicas=world_size, shuffle=False
)
)
sampler = (LengthGroupedSampler(
args.train_batch_size,
rank=rank,
world_size=world_size,
lengths=train_dataset.lengths,
group_frame=args.group_frame,
group_resolution=args.group_resolution,
) if (args.group_frame or args.group_resolution) else DistributedSampler(
train_dataset, rank=rank, num_replicas=world_size, shuffle=False))
train_dataloader = DataLoader(
train_dataset,
@@ -463,11 +388,7 @@ def main(args):
)
num_update_steps_per_epoch = math.ceil(
len(train_dataloader)
/ args.gradient_accumulation_steps
* args.sp_size
/ args.train_sp_batch_size
)
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
if rank <= 0:
@@ -475,22 +396,14 @@ def main(args):
wandb.init(project=project, config=args)
# Train!
total_batch_size = (
args.train_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" 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(
@@ -573,14 +486,12 @@ def main(args):
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix(
{
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
"phases": num_phases,
}
)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
"phases": num_phases,
})
progress_bar.update(1)
if rank <= 0:
wandb.log(
@@ -590,7 +501,7 @@ def main(args):
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
"pred_fro_norm": pred_norm["fro"],
"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"],
@@ -600,9 +511,7 @@ def main(args):
if step % args.checkpointing_steps == 0:
if args.use_lora:
# Save LoRA weights
save_lora_checkpoint(
transformer, optimizer, rank, args.output_dir, step
)
save_lora_checkpoint(transformer, optimizer, rank, args.output_dir, step)
else:
# Your existing checkpoint saving code
if args.use_ema:
@@ -640,9 +549,7 @@ 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)
@@ -653,12 +560,12 @@ def main(args):
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument(
"--model_type", type=str, default="mochi", help="The type of model to train."
)
parser.add_argument("--model_type", type=str, default="mochi", help="The type of model to train.")
# dataset & dataloader
parser.add_argument("--data_json_path", type=str, required=True)
parser.add_argument("--num_height", type=int, default=480)
parser.add_argument("--num_width", type=int, default=848)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
@@ -672,9 +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
@@ -696,9 +601,7 @@ 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,
@@ -715,39 +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
@@ -782,9 +677,7 @@ 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",
@@ -794,10 +687,8 @@ if __name__ == "__main__":
parser.add_argument(
"--allow_tf32",
action="store_true",
help=(
"Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
help=("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"),
)
parser.add_argument(
"--mixed_precision",
@@ -807,8 +698,7 @@ if __name__ == "__main__":
help=(
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
),
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."),
)
parser.add_argument(
"--use_cpu_offload",
@@ -830,12 +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
@@ -843,10 +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(
@@ -866,13 +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,
@@ -885,9 +765,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("--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)
+15 -38
View File
@@ -1,45 +1,23 @@
from typing import Any, Dict, Optional, Union
import torch
import torch.nn as nn
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
from diffusers.models.attention import JointTransformerBlock
from diffusers.models.attention_processor import Attention, AttentionProcessor
from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.normalization import AdaLayerNormContinuous
from diffusers.utils import (
USE_PEFT_BACKEND,
is_torch_version,
logging,
scale_lora_layers,
unscale_lora_layers,
)
from diffusers.models.embeddings import CombinedTimestepTextProjEmbeddings, PatchEmbed
from diffusers.models.transformers.transformer_2d import Transformer2DModelOutput
from diffusers.models.transformers.transformer_sd3 import SD3Transformer2DModel
from diffusers.utils import logging
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
class DiscriminatorHead(nn.Module):
def __init__(self, input_channel, output_channel=1):
super().__init__()
inner_channel = 1024
self.conv1 = nn.Sequential(
nn.Conv2d(input_channel, inner_channel, 1, 1, 0),
nn.GroupNorm(32, inner_channel),
nn.LeakyReLU(
inplace=True
), # use LeakyReLu instead of GELU shown in the paper to save memory
nn.LeakyReLU(inplace=True), # use LeakyReLu instead of GELU shown in the paper to save memory
)
self.conv2 = nn.Sequential(
nn.Conv2d(inner_channel, inner_channel, 1, 1, 0),
nn.GroupNorm(32, inner_channel),
nn.LeakyReLU(
inplace=True
), # use LeakyReLu instead of GELU shown in the paper to save memory
nn.LeakyReLU(inplace=True), # use LeakyReLu instead of GELU shown in the paper to save memory
)
self.conv_out = nn.Conv2d(inner_channel, output_channel, 1, 1, 0)
@@ -57,30 +35,29 @@ class DiscriminatorHead(nn.Module):
class Discriminator(nn.Module):
def __init__(
self, stride=8, num_h_per_head=1, adapter_channel_dims=[3072], total_layers=48,
self,
stride=8,
num_h_per_head=1,
adapter_channel_dims=[3072],
total_layers=48,
):
super().__init__()
adapter_channel_dims = adapter_channel_dims * (total_layers // stride)
self.stride = stride
self.num_h_per_head = num_h_per_head
self.head_num = len(adapter_channel_dims)
self.heads = nn.ModuleList(
[
nn.ModuleList(
[
DiscriminatorHead(adapter_channel)
for _ in range(self.num_h_per_head)
]
)
for adapter_channel in adapter_channel_dims
]
)
self.heads = nn.ModuleList([
nn.ModuleList([DiscriminatorHead(adapter_channel) for _ in range(self.num_h_per_head)])
for adapter_channel in adapter_channel_dims
])
def forward(self, features):
outputs = []
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs)
+31 -57
View File
@@ -3,11 +3,10 @@ from typing import Optional, Tuple, Union
import numpy as np
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.utils import BaseOutput, logging
from diffusers.utils.torch_utils import randn_tensor
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput, logging
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -21,7 +20,7 @@ class PCMFMSchedulerOutput(BaseOutput):
def extract_into_tensor(a, t, x_shape):
b, *_ = t.shape
out = a.gather(-1, t)
return out.reshape(b, *((1,) * (len(x_shape) - 1)))
return out.reshape(b, *((1, ) * (len(x_shape) - 1)))
class PCMFMScheduler(SchedulerMixin, ConfigMixin):
@@ -40,20 +39,15 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
):
if linear_quadratic:
linear_steps = int(num_train_timesteps * linear_range)
sigmas = linear_quadratic_schedule(
num_train_timesteps, linear_quadratic_threshold, linear_steps
)
sigmas = linear_quadratic_schedule(num_train_timesteps, linear_quadratic_threshold, linear_steps)
sigmas = torch.tensor(sigmas).to(dtype=torch.float32)
else:
timesteps = np.linspace(
1, num_train_timesteps, num_train_timesteps, dtype=np.float32
)[::-1].copy()
timesteps = np.linspace(1, num_train_timesteps, num_train_timesteps, dtype=np.float32)[::-1].copy()
timesteps = torch.from_numpy(timesteps).to(dtype=torch.float32)
sigmas = timesteps / num_train_timesteps
sigmas = shift * sigmas / (1 + (shift - 1) * sigmas)
self.euler_timesteps = (
np.arange(1, pcm_timesteps + 1) * (num_train_timesteps // pcm_timesteps)
).round().astype(np.int64) - 1
self.euler_timesteps = (np.arange(1, pcm_timesteps + 1) *
(num_train_timesteps // pcm_timesteps)).round().astype(np.int64) - 1
self.sigmas = sigmas.numpy()[::-1][self.euler_timesteps]
self.sigmas = torch.from_numpy((self.sigmas[::-1].copy()))
self.timesteps = self.sigmas * num_train_timesteps
@@ -118,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).
@@ -131,18 +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
@@ -204,18 +192,11 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
returned, otherwise a tuple is returned where the first element is the sample tensor.
"""
if (
isinstance(timestep, int)
or isinstance(timestep, torch.IntTensor)
or isinstance(timestep, torch.LongTensor)
):
raise ValueError(
(
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."
),
)
if (isinstance(timestep, int) or isinstance(timestep, torch.IntTensor)
or isinstance(timestep, torch.LongTensor)):
raise ValueError(("Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."), )
if self.step_index is None:
self._init_step_index(timestep)
@@ -233,7 +214,7 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
self._step_index += 1
if not return_dict:
return (prev_sample,)
return (prev_sample, )
return PCMFMSchedulerOutput(prev_sample=prev_sample)
@@ -242,16 +223,14 @@ class PCMFMScheduler(SchedulerMixin, ConfigMixin):
class EulerSolver:
def __init__(self, sigmas, timesteps=1000, euler_timesteps=50):
self.step_ratio = timesteps // euler_timesteps
self.euler_timesteps = (
np.arange(1, euler_timesteps + 1) * self.step_ratio
).round().astype(np.int64) - 1
self.euler_timesteps = (np.arange(1, euler_timesteps + 1) * self.step_ratio).round().astype(np.int64) - 1
self.euler_timesteps_prev = np.asarray([0] + self.euler_timesteps[:-1].tolist())
self.sigmas = sigmas[self.euler_timesteps]
self.sigmas_prev = np.asarray(
[sigmas[0]] + sigmas[self.euler_timesteps[:-1]].tolist()
) # either use sigma0 or 0
self.sigmas_prev = np.asarray([sigmas[0]] +
sigmas[self.euler_timesteps[:-1]].tolist()) # either use sigma0 or 0
self.euler_timesteps = torch.from_numpy(self.euler_timesteps).long()
self.euler_timesteps_prev = torch.from_numpy(self.euler_timesteps_prev).long()
@@ -268,25 +247,22 @@ class EulerSolver:
def euler_step(self, sample, model_pred, timestep_index):
sigma = extract_into_tensor(self.sigmas, timestep_index, model_pred.shape)
sigma_prev = extract_into_tensor(
self.sigmas_prev, timestep_index, model_pred.shape
)
sigma_prev = extract_into_tensor(self.sigmas_prev, timestep_index, model_pred.shape)
x_prev = sample + (sigma_prev - sigma) * model_pred
return x_prev
def euler_style_multiphase_pred(
self, sample, model_pred, timestep_index, multiphase, is_target=False,
self,
sample,
model_pred,
timestep_index,
multiphase,
is_target=False,
):
inference_indices = np.linspace(
0, len(self.euler_timesteps), num=multiphase, endpoint=False
)
inference_indices = np.linspace(0, len(self.euler_timesteps), num=multiphase, endpoint=False)
inference_indices = np.floor(inference_indices).astype(np.int64)
inference_indices = (
torch.from_numpy(inference_indices).long().to(self.euler_timesteps.device)
)
expanded_timestep_index = timestep_index.unsqueeze(1).expand(
-1, inference_indices.size(0)
)
inference_indices = (torch.from_numpy(inference_indices).long().to(self.euler_timesteps.device))
expanded_timestep_index = timestep_index.unsqueeze(1).expand(-1, inference_indices.size(0))
valid_indices_mask = expanded_timestep_index >= inference_indices
last_valid_index = valid_indices_mask.flip(dims=[1]).long().argmax(dim=1)
last_valid_index = inference_indices.size(0) - 1 - last_valid_index
@@ -296,9 +272,7 @@ class EulerSolver:
sigma = extract_into_tensor(self.sigmas_prev, timestep_index, sample.shape)
else:
sigma = extract_into_tensor(self.sigmas, timestep_index, sample.shape)
sigma_prev = extract_into_tensor(
self.sigmas_prev, timestep_index_end, sample.shape
)
sigma_prev = extract_into_tensor(self.sigmas_prev, timestep_index_end, sample.shape)
x_prev = sample + (sigma_prev - sigma) * model_pred
return x_prev, timestep_index_end
+114 -201
View File
@@ -1,67 +1,43 @@
# !/bin/python3
# isort: skip_file
import argparse
from email.policy import strict
import logging
import math
import os
import shutil
from pathlib import Path
from fastvideo.utils.parallel_states import (
initialize_sequence_parallel_state,
destroy_sequence_parallel_group,
get_sequence_parallel_state,
nccl_info,
)
from fastvideo.utils.communications import sp_parallel_dataloader_wrapper, broadcast
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.utils.validation import log_validation
import time
from torch.utils.data import DataLoader
import torch
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
StateDictType,
FullStateDictConfig,
)
from fastvideo.utils.load import load_transformer
from collections import deque
from copy import deepcopy
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
import json
from torch.utils.data.distributed import DistributedSampler
from fastvideo.utils.dataset_utils import LengthGroupedSampler
import torch
import torch.distributed as dist
import wandb
from accelerate.utils import set_seed
from tqdm.auto import tqdm
from fastvideo.utils.fsdp_util import (
get_dit_fsdp_kwargs,
apply_fsdp_checkpointing,
get_discriminator_fsdp_kwargs,
)
import diffusers
from diffusers import FlowMatchEulerDiscreteScheduler
from fastvideo.distill.discriminator import Discriminator
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from copy import deepcopy
from diffusers.optimization import get_scheduler
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
from diffusers.utils import check_min_version
from fastvideo.dataset.latent_datasets import LatentDataset, latent_collate_function
import torch.distributed as dist
from peft import LoraConfig
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
from fastvideo.utils.checkpoint import (
save_checkpoint,
save_lora_checkpoint,
resume_lora_optimizer,
resume_training,
save_checkpoint_generator_discriminator,
resume_training_generator_discriminator,
)
from torch.utils.data import DataLoader
from torch.utils.data.distributed import DistributedSampler
from tqdm.auto import tqdm
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
from fastvideo.distill.discriminator import Discriminator
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
from fastvideo.utils.checkpoint import (resume_lora_optimizer, resume_training_generator_discriminator, save_checkpoint,
save_lora_checkpoint)
from fastvideo.utils.communications import (broadcast, sp_parallel_dataloader_wrapper)
from fastvideo.utils.dataset_utils import LengthGroupedSampler
from fastvideo.utils.fsdp_util import (apply_fsdp_checkpointing, get_discriminator_fsdp_kwargs, get_dit_fsdp_kwargs)
from fastvideo.utils.load import load_transformer
from fastvideo.utils.logging_ import main_print
from fastvideo.utils.parallel_states import (destroy_sequence_parallel_group, get_sequence_parallel_state,
initialize_sequence_parallel_state)
from fastvideo.utils.validation import log_validation
# Will error if the minimal version of diffusers is not installed. Remove at your own risks.
check_min_version("0.31.0")
import time
from collections import deque
def gan_d_loss(
@@ -100,10 +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
@@ -127,11 +101,10 @@ def gan_g_loss(
output_features_stride=discriminator_head_stride,
return_dict=False,
)[1]
fake_outputs = discriminator(features,)
fake_outputs = discriminator(features, )
for fake_output in fake_outputs:
loss += torch.mean(weight * torch.relu(1 - fake_output.float())) / (
discriminator.head_num * discriminator.num_h_per_head
)
loss += torch.mean(
weight * torch.relu(1 - fake_output.float())) / (discriminator.head_num * discriminator.num_h_per_head)
return loss
@@ -170,9 +143,7 @@ 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.
@@ -197,11 +168,8 @@ def distill_one_step_adv(
)[0]
# if accelerator.is_main_process:
model_pred, end_index = solver.euler_style_multiphase_pred(
noisy_model_input, model_pred, index, multiphase
)
model_pred, end_index = solver.euler_style_multiphase_pred(noisy_model_input, model_pred, index, multiphase)
weighting = 1.0
# # simplified flow matching aka 0-rectified flow matching loss
# # target = model_input - noise
# target = model_input
@@ -210,14 +178,13 @@ def distill_one_step_adv(
adv_index[i] = torch.randint(
end_index[i].item(),
end_index[i].item() + num_euler_timesteps // multiphase,
(1,),
(1, ),
dtype=end_index.dtype,
device=end_index.device,
)
sigmas_end = extract_into_tensor(solver.sigmas_prev, end_index, model_input.shape)
sigmas_adv = extract_into_tensor(solver.sigmas_prev, adv_index, model_input.shape)
timesteps_end = (sigmas_end * noise_scheduler.config.num_train_timesteps).view(-1)
timesteps_adv = (sigmas_adv * noise_scheduler.config.num_train_timesteps).view(-1)
with torch.no_grad():
@@ -242,9 +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
@@ -258,22 +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)
fake_adv = (
(1 - sigmas_adv) * model_pred
+ (sigmas_adv - sigmas_end) * torch.randn_like(model_pred)
) / (1 - sigmas_end)
real_adv = ((1 - sigmas_adv) * target + (sigmas_adv - sigmas_end) * torch.randn_like(target)) / (1 - sigmas_end)
fake_adv = ((1 - sigmas_adv) * model_pred +
(sigmas_adv - sigmas_end) * torch.randn_like(model_pred)) / (1 - sigmas_end)
huber_c = 0.001
g_loss = torch.mean(
torch.sqrt((model_pred.float() - target.float()) ** 2 + huber_c ** 2) - huber_c
)
g_loss = torch.mean(torch.sqrt((model_pred.float() - target.float())**2 + huber_c**2) - huber_c)
discriminator.requires_grad_(False)
with torch.autocast("cuda", dtype=torch.bfloat16):
g_gan_loss = adv_weight * gan_g_loss(
@@ -342,7 +299,7 @@ def main(args):
if rank <= 0 and args.output_dir is not None:
os.makedirs(args.output_dir, exist_ok=True)
# For mixed precision training we cast all non-trainable weigths to half-precision
# For mixed precision training we cast all non-trainable weights to half-precision
# as these weights are only used for inference, keeping weights in full precision is not required.
# Create model:
@@ -378,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,
@@ -397,26 +352,29 @@ def main(args):
transformer._no_split_modules = no_split_modules
fsdp_kwargs["auto_wrap_policy"] = fsdp_kwargs["auto_wrap_policy"](transformer)
transformer = FSDP(transformer, **fsdp_kwargs,)
teacher_transformer = FSDP(teacher_transformer, **fsdp_kwargs,)
discriminator = FSDP(discriminator, **discriminator_fsdp_kwargs,)
main_print(f"--> model loaded")
transformer = FSDP(
transformer,
**fsdp_kwargs,
)
teacher_transformer = FSDP(
teacher_transformer,
**fsdp_kwargs,
)
discriminator = FSDP(
discriminator,
**discriminator_fsdp_kwargs,
)
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
@@ -433,7 +391,7 @@ def main(args):
params_to_optimize,
lr=args.learning_rate,
betas=(0.9, 0.999),
weight_decay=1e-3,
weight_decay=args.weight_decay,
eps=1e-8,
)
@@ -441,15 +399,14 @@ def main(args):
discriminator.parameters(),
lr=args.discriminator_learning_rate,
betas=(0, 0.999),
weight_decay=1e-3,
weight_decay=args.weight_decay,
eps=1e-8,
)
init_steps = 0
if args.resume_from_lora_checkpoint:
transformer, optimizer, init_steps = resume_lora_optimizer(
transformer, args.resume_from_lora_checkpoint, optimizer
)
transformer, optimizer, init_steps = resume_lora_optimizer(transformer, args.resume_from_lora_checkpoint,
optimizer)
elif args.resume_from_checkpoint:
(
transformer,
@@ -481,20 +438,15 @@ def main(args):
train_dataset = LatentDataset(args.data_json_path, args.num_latent_t, args.cfg)
uncond_prompt_embed = train_dataset.uncond_prompt_embed
uncond_prompt_mask = train_dataset.uncond_prompt_mask
sampler = (
LengthGroupedSampler(
args.train_batch_size,
rank=rank,
world_size=world_size,
lengths=train_dataset.lengths,
group_frame=args.group_frame,
group_resolution=args.group_resolution,
)
if (args.group_frame or args.group_resolution)
else DistributedSampler(
train_dataset, rank=rank, num_replicas=world_size, shuffle=False
)
)
sampler = (LengthGroupedSampler(
args.train_batch_size,
rank=rank,
world_size=world_size,
lengths=train_dataset.lengths,
group_frame=args.group_frame,
group_resolution=args.group_resolution,
) if (args.group_frame or args.group_resolution) else DistributedSampler(
train_dataset, rank=rank, num_replicas=world_size, shuffle=False))
train_dataloader = DataLoader(
train_dataset,
@@ -507,11 +459,7 @@ def main(args):
)
assert args.gradient_accumulation_steps == 1
num_update_steps_per_epoch = math.ceil(
len(train_dataloader)
/ args.gradient_accumulation_steps
* args.sp_size
/ args.train_sp_batch_size
)
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
if rank <= 0:
@@ -519,22 +467,14 @@ def main(args):
wandb.init(project=project, config=args)
# Train!
total_batch_size = (
args.train_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" 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(
@@ -559,6 +499,7 @@ def main(args):
)
step_times = deque(maxlen=100)
# log_validation(args, transformer, device,
# torch.bfloat16, 0, scheduler_type=args.scheduler_type, shift=args.shift, num_euler_timesteps=args.num_euler_timesteps, linear_quadratic_threshold=args.linear_quadratic_threshold,ema=False)
def get_num_phases(multi_phased_distill_schedule, step):
@@ -610,15 +551,13 @@ def main(args):
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix(
{
"g_loss": f"{generator_loss:.4f}",
"d_loss": f"{discriminator_loss:.4f}",
"g_grad_norm": generator_grad_norm,
"d_grad_norm": discriminator_grad_norm,
"step_time": f"{step_time:.2f}s",
}
)
progress_bar.set_postfix({
"g_loss": f"{generator_loss:.4f}",
"d_loss": f"{discriminator_loss:.4f}",
"g_grad_norm": generator_grad_norm,
"d_grad_norm": discriminator_grad_norm,
"step_time": f"{step_time:.2f}s",
})
progress_bar.update(1)
if rank <= 0:
wandb.log(
@@ -637,9 +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,9 +589,7 @@ def main(args):
# args.output_dir,
# step,
# )
save_checkpoint(
transformer, rank, args.output_dir, args.max_train_steps
)
save_checkpoint(transformer, rank, args.output_dir, step)
main_print(f"--> checkpoint saved at step {step}")
dist.barrier()
if args.log_validation and step % args.validation_steps == 0:
@@ -673,9 +608,7 @@ 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)
@@ -686,11 +619,11 @@ def main(args):
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument(
"--model_type", type=str, default="mochi", help="The type of model to train."
)
parser.add_argument("--model_type", type=str, default="mochi", help="The type of model to train.")
# dataset & dataloader
parser.add_argument("--data_json_path", type=str, required=True)
parser.add_argument("--num_height", type=int, default=480)
parser.add_argument("--num_width", type=int, default=848)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
@@ -704,9 +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
@@ -725,9 +656,7 @@ 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,
@@ -744,11 +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)
@@ -756,28 +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
@@ -812,9 +733,7 @@ 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",
@@ -824,10 +743,8 @@ if __name__ == "__main__":
parser.add_argument(
"--allow_tf32",
action="store_true",
help=(
"Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"
),
help=("Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see"
" https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices"),
)
parser.add_argument(
"--mixed_precision",
@@ -837,8 +754,7 @@ if __name__ == "__main__":
help=(
"Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >="
" 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the"
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."
),
" flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config."),
)
parser.add_argument(
"--use_cpu_offload",
@@ -860,12 +776,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")
parser.add_argument("--multi_phased_distill_schedule", type=str, default=None)
parser.add_argument(
@@ -880,10 +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(
@@ -903,13 +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,
@@ -922,6 +828,13 @@ if __name__ == "__main__":
default=2,
help="The stride of the discriminator head.",
)
parser.add_argument(
"--linear_range",
type=float,
default=0.5,
help="Range for linear quadratic scheduler.",
)
parser.add_argument("--weight_decay", type=float, default=0.001, help="Weight decay to apply.")
parser.add_argument(
"--linear_quadratic_threshold",
type=float,
+4 -10
View File
@@ -1,19 +1,15 @@
from einops import rearrange
from flash_attn import flash_attn_varlen_qkvpacked_func
from flash_attn.bert_padding import pad_input, unpad_input
from einops import rearrange
def flash_attn_no_pad(
qkv, key_padding_mask, causal=False, dropout_p=0.0, softmax_scale=None
):
def flash_attn_no_pad(qkv, key_padding_mask, causal=False, dropout_p=0.0, softmax_scale=None):
# adapted from https://github.com/Dao-AILab/flash-attention/blob/13403e81157ba37ca525890f2f0f2137edf75311/flash_attn/flash_attention.py#L27
batch_size = qkv.shape[0]
seqlen = qkv.shape[1]
nheads = qkv.shape[-2]
x = rearrange(qkv, "b s three h d -> b s (three h d)")
x_unpad, indices, cu_seqlens, max_s, used_seqlens_in_batch = unpad_input(
x, key_padding_mask
)
x_unpad, indices, cu_seqlens, max_s, used_seqlens_in_batch = unpad_input(x, key_padding_mask)
x_unpad = rearrange(x_unpad, "nnz (three h d) -> nnz three h d", three=3, h=nheads)
output_unpad = flash_attn_varlen_qkvpacked_func(
@@ -25,9 +21,7 @@ def flash_attn_no_pad(
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,
)
+7 -5
View File
@@ -1,4 +1,5 @@
import os
import torch
__all__ = [
@@ -33,8 +34,7 @@ C_SCALE = 1_000_000_000_000_000
PROMPT_TEMPLATE_ENCODE = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the image by detailing the color, shape, size, texture, "
"quantity, text, spatial relationships of the objects and background:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
)
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
PROMPT_TEMPLATE_ENCODE_VIDEO = (
"<|start_header_id|>system<|end_header_id|>\n\nDescribe the video by detailing the following aspects: "
"1. The main content and theme of the video."
@@ -42,13 +42,15 @@ PROMPT_TEMPLATE_ENCODE_VIDEO = (
"3. Actions, events, behaviors temporal relationships, physical movement changes of the objects."
"4. background environment, light, style and atmosphere."
"5. camera angles, movements, and transitions used in the video:<|eot_id|>"
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>"
)
"<|start_header_id|>user<|end_header_id|>\n\n{}<|eot_id|>")
NEGATIVE_PROMPT = "Aerial view, aerial view, overexposed, low quality, deformation, a poor composition, bad hands, bad teeth, bad eyes, bad limbs, distortion"
PROMPT_TEMPLATE = {
"dit-llm-encode": {"template": PROMPT_TEMPLATE_ENCODE, "crop_start": 36,},
"dit-llm-encode": {
"template": PROMPT_TEMPLATE_ENCODE,
"crop_start": 36,
},
"dit-llm-encode-video": {
"template": PROMPT_TEMPLATE_ENCODE_VIDEO,
"crop_start": 95,
@@ -1,2 +1,3 @@
# ruff: noqa: F401
from .pipelines import HunyuanVideoPipeline
from .schedulers import FlowMatchDiscreteScheduler
@@ -1 +1,2 @@
# ruff: noqa: F401
from .pipeline_hunyuan_video import HunyuanVideoPipeline
@@ -17,41 +17,33 @@
#
# ==============================================================================
import inspect
from typing import Any, Callable, Dict, List, Optional, Union, Tuple
from dataclasses import dataclass
from typing import Any, Callable, Dict, List, Optional, Union
import numpy as np
import torch
import torch.distributed as dist
import numpy as np
from dataclasses import dataclass
from packaging import version
import torch.nn.functional as F
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
from diffusers.configuration_utils import FrozenDict
from diffusers.image_processor import VaeImageProcessor
from diffusers.loaders import LoraLoaderMixin, TextualInversionLoaderMixin
from diffusers.models import AutoencoderKL
from diffusers.models.lora import adjust_lora_scale_text_encoder
from diffusers.schedulers import KarrasDiffusionSchedulers
from diffusers.utils import (
USE_PEFT_BACKEND,
deprecate,
logging,
replace_example_docstring,
scale_lora_layers,
unscale_lora_layers,
)
from diffusers.utils.torch_utils import randn_tensor
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
from diffusers.utils import BaseOutput
from diffusers.schedulers import KarrasDiffusionSchedulers
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 ...constants import PRECISION_TO_TYPE
from ...vae.autoencoder_kl_causal_3d import AutoencoderKLCausal3D
from ...text_encoder import TextEncoder
from ...modules import HYVideoDiffusionTransformer
from einops import rearrange
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
from fastvideo.utils.communications import all_gather, all_to_all_4D
import torch.nn.functional as F
from ...text_encoder import TextEncoder
from ...vae.autoencoder_kl_causal_3d import AutoencoderKLCausal3D
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -63,16 +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
@@ -108,30 +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)
@@ -193,39 +173,27 @@ class HunyuanVideoPipeline(DiffusionPipeline):
self.args = args
# ==========================================================================================
if (
hasattr(scheduler.config, "steps_offset")
and scheduler.config.steps_offset != 1
):
if (hasattr(scheduler.config, "steps_offset") and scheduler.config.steps_offset != 1):
deprecation_message = (
f"The configuration file of this scheduler: {scheduler} is outdated. `steps_offset`"
f" should be set to 1 instead of {scheduler.config.steps_offset}. Please make sure "
"to update the config accordingly as leaving `steps_offset` might led to incorrect results"
" in future versions. If you have downloaded this checkpoint from the Hugging Face Hub,"
" it would be very nice if you could open a Pull request for the `scheduler/scheduler_config.json`"
" file"
)
deprecate(
"steps_offset!=1", "1.0.0", deprecation_message, standard_warn=False
)
" file")
deprecate("steps_offset!=1", "1.0.0", deprecation_message, standard_warn=False)
new_config = dict(scheduler.config)
new_config["steps_offset"] = 1
scheduler._internal_dict = FrozenDict(new_config)
if (
hasattr(scheduler.config, "clip_sample")
and scheduler.config.clip_sample is True
):
if (hasattr(scheduler.config, "clip_sample") and scheduler.config.clip_sample is True):
deprecation_message = (
f"The configuration file of this scheduler: {scheduler} has not set the configuration `clip_sample`."
" `clip_sample` should be set to False in the configuration file. Please make sure to update the"
" config accordingly as not setting `clip_sample` in the config might lead to incorrect results in"
" future versions. If you have downloaded this checkpoint from the Hugging Face Hub, it would be very"
" nice if you could open a Pull request for the `scheduler/scheduler_config.json` file"
)
deprecate(
"clip_sample not set", "1.0.0", deprecation_message, standard_warn=False
)
" nice if you could open a Pull request for the `scheduler/scheduler_config.json` file")
deprecate("clip_sample not set", "1.0.0", deprecation_message, standard_warn=False)
new_config = dict(scheduler.config)
new_config["clip_sample"] = False
scheduler._internal_dict = FrozenDict(new_config)
@@ -237,7 +205,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
scheduler=scheduler,
text_encoder_2=text_encoder_2,
)
self.vae_scale_factor = 2 ** (len(self.vae.config.block_out_channels) - 1)
self.vae_scale_factor = 2**(len(self.vae.config.block_out_channels) - 1)
self.image_processor = VaeImageProcessor(vae_scale_factor=self.vae_scale_factor)
def encode_prompt(
@@ -303,13 +271,6 @@ class HunyuanVideoPipeline(DiffusionPipeline):
else:
scale_lora_layers(text_encoder.model, lora_scale)
if prompt is not None and isinstance(prompt, str):
batch_size = 1
elif prompt is not None and isinstance(prompt, list):
batch_size = len(prompt)
else:
batch_size = prompt_embeds.shape[0]
if prompt_embeds is None:
# textual inversion: process multi-vector tokens if necessary
if isinstance(self, TextualInversionLoaderMixin):
@@ -317,9 +278,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
text_inputs = text_encoder.text2tokens(prompt, data_type=data_type)
if clip_skip is None:
prompt_outputs = text_encoder.encode(
text_inputs, data_type=data_type, device=device
)
prompt_outputs = text_encoder.encode(text_inputs, data_type=data_type, device=device)
prompt_embeds = prompt_outputs.hidden_state
else:
prompt_outputs = text_encoder.encode(
@@ -336,18 +295,14 @@ class HunyuanVideoPipeline(DiffusionPipeline):
# representations. The `last_hidden_states` that we typically use for
# obtaining the final prompt representations passes through the LayerNorm
# layer.
prompt_embeds = text_encoder.model.text_model.final_layer_norm(
prompt_embeds
)
prompt_embeds = text_encoder.model.text_model.final_layer_norm(prompt_embeds)
attention_mask = prompt_outputs.attention_mask
if attention_mask is not None:
attention_mask = attention_mask.to(device)
bs_embed, seq_len = attention_mask.shape
attention_mask = attention_mask.repeat(1, num_videos_per_prompt)
attention_mask = attention_mask.view(
bs_embed * num_videos_per_prompt, seq_len
)
attention_mask = attention_mask.view(bs_embed * num_videos_per_prompt, seq_len)
if text_encoder is not None:
prompt_embeds_dtype = text_encoder.dtype
@@ -367,9 +322,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
bs_embed, seq_len, _ = prompt_embeds.shape
# duplicate text embeddings for each generation per prompt, using mps friendly method
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
prompt_embeds = prompt_embeds.view(
bs_embed * num_videos_per_prompt, seq_len, -1
)
prompt_embeds = prompt_embeds.view(bs_embed * num_videos_per_prompt, seq_len, -1)
return (
prompt_embeds,
@@ -385,9 +338,7 @@ class HunyuanVideoPipeline(DiffusionPipeline):
latents = 1 / self.vae.config.scaling_factor * latents
if enable_tiling:
self.vae.enable_tiling()
image = self.vae.decode(latents, return_dict=False)[0]
else:
image = self.vae.decode(latents, return_dict=False)[0]
image = self.vae.decode(latents, return_dict=False)[0]
image = (image / 2 + 0.5).clamp(0, 1)
# we always cast to float32 as this does not cause significant overhead and is compatible with bfloat16
if image.ndim == 4:
@@ -423,33 +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]}"
)
@@ -457,32 +396,23 @@ class HunyuanVideoPipeline(DiffusionPipeline):
if prompt is not None and prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
" only forward one of the two."
)
" only forward one of the two.")
elif prompt is None and prompt_embeds is None:
raise ValueError(
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
)
elif prompt is not None and (
not isinstance(prompt, str) and not isinstance(prompt, list)
):
raise ValueError(
f"`prompt` has to be of type `str` or `list` but is {type(prompt)}"
)
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined.")
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
if negative_prompt is not None and negative_prompt_embeds is not None:
raise ValueError(
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
f" {negative_prompt_embeds}. Please make sure to only forward one of the two."
)
raise ValueError(f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`:"
f" {negative_prompt_embeds}. Please make sure to only forward one of the two.")
if prompt_embeds is not None and negative_prompt_embeds is not None:
if prompt_embeds.shape != negative_prompt_embeds.shape:
raise ValueError(
"`prompt_embeds` and `negative_prompt_embeds` must have the same shape when passed directly, but"
f" got: `prompt_embeds` {prompt_embeds.shape} != `negative_prompt_embeds`"
f" {negative_prompt_embeds.shape}."
)
f" {negative_prompt_embeds.shape}.")
def prepare_latents(
self,
@@ -506,13 +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)
@@ -615,18 +542,15 @@ class HunyuanVideoPipeline(DiffusionPipeline):
cross_attention_kwargs: Optional[Dict[str, Any]] = None,
guidance_rescale: float = 0.0,
clip_skip: Optional[int] = None,
callback_on_step_end: Optional[
Union[
Callable[[int, int, Dict], None],
PipelineCallback,
MultiPipelineCallbacks,
]
] = None,
callback_on_step_end: Optional[Union[Callable[[int, int, Dict], None], PipelineCallback,
MultiPipelineCallbacks, ]] = None,
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
vae_ver: str = "88-4c-sd",
enable_tiling: bool = False,
enable_vae_sp: bool = False,
n_tokens: Optional[int] = None,
embedded_guidance_scale: Optional[float] = None,
mask_strategy: Optional[Dict[str, list]] = None,
**kwargs,
):
r"""
@@ -763,18 +687,11 @@ class HunyuanVideoPipeline(DiffusionPipeline):
else:
batch_size = prompt_embeds.shape[0]
device = (
torch.device(f"cuda:{dist.get_rank()}")
if dist.is_initialized()
else self._execution_device
)
device = (torch.device(f"cuda:{dist.get_rank()}") if dist.is_initialized() else self._execution_device)
# 3. Encode input prompt
lora_scale = (
self.cross_attention_kwargs.get("scale", None)
if self.cross_attention_kwargs is not None
else None
)
lora_scale = (self.cross_attention_kwargs.get("scale", None)
if self.cross_attention_kwargs is not None else None)
(
prompt_embeds,
@@ -835,9 +752,8 @@ class HunyuanVideoPipeline(DiffusionPipeline):
prompt_mask_2 = torch.cat([negative_prompt_mask_2, prompt_mask_2])
# 4. Prepare timesteps
extra_set_timesteps_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.set_timesteps, {"n_tokens": n_tokens}
)
extra_set_timesteps_kwargs = self.prepare_extra_func_kwargs(self.scheduler.set_timesteps,
{"n_tokens": n_tokens})
timesteps, num_inference_steps = retrieve_timesteps(
self.scheduler,
num_inference_steps,
@@ -866,32 +782,39 @@ class HunyuanVideoPipeline(DiffusionPipeline):
generator,
latents,
)
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
if get_sequence_parallel_state():
latents = rearrange(
latents, "b t (n s) h w -> b t n s h w", n=world_size
).contiguous()
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
latents = latents[:, :, rank, :, :, :]
# 6. Prepare extra step kwargs. TODO: Logic should ideally just be moved out of the pipeline
extra_step_kwargs = self.prepare_extra_func_kwargs(
self.scheduler.step, {"generator": generator, "eta": eta},
self.scheduler.step,
{
"generator": generator,
"eta": eta
},
)
target_dtype = PRECISION_TO_TYPE[self.args.precision]
autocast_enabled = (
target_dtype != torch.float32
) and not self.args.disable_autocast
autocast_enabled = (target_dtype != torch.float32) and not self.args.disable_autocast
vae_dtype = PRECISION_TO_TYPE[self.args.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32
) and not self.args.disable_autocast
vae_autocast_enabled = (vae_dtype != torch.float32) and not self.args.disable_autocast
# 7. Denoising loop
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
self._num_timesteps = len(timesteps)
def dict_to_3d_list(mask_strategy, t_max=50, l_max=60, h_max=24):
result = [[[None for _ in range(h_max)] for _ in range(l_max)] for _ in range(t_max)]
if mask_strategy is None:
return result
for key, value in mask_strategy.items():
t, l, h = map(int, key.split('_'))
result[t][l][h] = value
return result
mask_strategy = dict_to_3d_list(mask_strategy)
# if is_progress_bar:
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
@@ -899,57 +822,39 @@ class HunyuanVideoPipeline(DiffusionPipeline):
continue
# expand the latents if we are doing classifier free guidance
latent_model_input = (
torch.cat([latents] * 2)
if self.do_classifier_free_guidance
else latents
)
latent_model_input = self.scheduler.scale_model_input(
latent_model_input, t
)
latent_model_input = (torch.cat([latents] * 2) if self.do_classifier_free_guidance else latents)
latent_model_input = self.scheduler.scale_model_input(latent_model_input, t)
t_expand = t.repeat(latent_model_input.shape[0])
guidance_expand = (
torch.tensor(
[embedded_guidance_scale] * latent_model_input.shape[0],
dtype=torch.float32,
device=device,
).to(target_dtype)
* 1000.0
if embedded_guidance_scale is not None
else None
)
guidance_expand = (torch.tensor(
[embedded_guidance_scale] * latent_model_input.shape[0],
dtype=torch.float32,
device=device,
).to(target_dtype) * 1000.0 if embedded_guidance_scale is not None else None)
# predict the noise residual
with torch.autocast(
device_type="cuda", dtype=target_dtype, enabled=autocast_enabled
):
# concat prompt_embeds_2 and prompt_embeds. Mismach fill with zeros
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]),
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
]
)[0]
# perform guidance
if self.do_classifier_free_guidance:
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
noise_pred = noise_pred_uncond + self.guidance_scale * (
noise_pred_text - noise_pred_uncond
)
noise_pred = noise_pred_uncond + self.guidance_scale * (noise_pred_text - noise_pred_uncond)
if self.do_classifier_free_guidance and self.guidance_rescale > 0.0:
# Based on 3.4. in https://arxiv.org/pdf/2305.08891.pdf
@@ -960,9 +865,7 @@ 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 = {}
@@ -972,14 +875,10 @@ class HunyuanVideoPipeline(DiffusionPipeline):
latents = callback_outputs.pop("latents", latents)
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
negative_prompt_embeds = callback_outputs.pop(
"negative_prompt_embeds", negative_prompt_embeds
)
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
# call the callback, if provided
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0
):
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
if progress_bar is not None:
progress_bar.update()
if callback is not None and i % callback_steps == 0:
@@ -999,32 +898,19 @@ class HunyuanVideoPipeline(DiffusionPipeline):
pass
else:
raise ValueError(
f"Only support latents with shape (b, c, h, w) or (b, c, f, h, w), but got {latents.shape}."
)
f"Only support latents with shape (b, c, h, w) or (b, c, f, h, w), but got {latents.shape}.")
if (
hasattr(self.vae.config, "shift_factor")
and self.vae.config.shift_factor
):
latents = (
latents / self.vae.config.scaling_factor
+ self.vae.config.shift_factor
)
if (hasattr(self.vae.config, "shift_factor") and self.vae.config.shift_factor):
latents = (latents / self.vae.config.scaling_factor + self.vae.config.shift_factor)
else:
latents = latents / self.vae.config.scaling_factor
with torch.autocast(
device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled
):
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
if enable_tiling:
self.vae.enable_tiling()
image = self.vae.decode(
latents, return_dict=False, generator=generator
)[0]
else:
image = self.vae.decode(
latents, return_dict=False, generator=generator
)[0]
if enable_vae_sp:
self.vae.enable_parallel()
image = self.vae.decode(latents, return_dict=False, generator=generator)[0]
if expand_temporal_dim or image.shape[2] == 1:
image = image.squeeze(2)
@@ -1 +1,2 @@
# ruff: noqa: F401
from .scheduling_flow_match_discrete import FlowMatchDiscreteScheduler
@@ -20,13 +20,10 @@
from dataclasses import dataclass
from typing import Optional, Tuple, Union
import numpy as np
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.utils import BaseOutput, logging
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput, logging
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
@@ -90,9 +87,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
self.supported_solver = ["euler"]
if solver not in self.supported_solver:
raise ValueError(
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
)
raise ValueError(f"Solver {solver} not supported. Supported solvers: {self.supported_solver}")
@property
def step_index(self):
@@ -148,9 +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
@@ -177,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):
@@ -217,18 +208,11 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
returned, otherwise a tuple is returned where the first element is the sample tensor.
"""
if (
isinstance(timestep, int)
or isinstance(timestep, torch.IntTensor)
or isinstance(timestep, torch.LongTensor)
):
raise ValueError(
(
"Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."
),
)
if (isinstance(timestep, int) or isinstance(timestep, torch.IntTensor)
or isinstance(timestep, torch.LongTensor)):
raise ValueError(("Passing integer indices (e.g. from `enumerate(timesteps)`) as timesteps to"
" `EulerDiscreteScheduler.step()` is not supported. Make sure to pass"
" one of the `scheduler.timesteps` as a timestep."), )
if self.step_index is None:
self._init_step_index(timestep)
@@ -241,15 +225,13 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
if self.config.solver == "euler":
prev_sample = sample + model_output.to(torch.float32) * dt
else:
raise ValueError(
f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}"
)
raise ValueError(f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}")
# upon completion increase step index by one
self._step_index += 1
if not return_dict:
return (prev_sample,)
return (prev_sample, )
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
+23 -26
View File
@@ -1,6 +1,8 @@
# ruff: noqa: F405, F403
import argparse
from .constants import *
import re
from .constants import *
from .modules.models import HUNYUAN_VIDEO_CONFIG
@@ -45,16 +47,12 @@ def add_network_args(parser: argparse.ArgumentParser):
)
# RoPE
group.add_argument(
"--rope-theta", type=int, default=256, help="Theta used in RoPE."
)
group.add_argument("--rope-theta", type=int, default=256, help="Theta used in RoPE.")
return parser
def add_extra_models_args(parser: argparse.ArgumentParser):
group = parser.add_argument_group(
title="Extra models args, including vae, text encoders and tokenizers)"
)
group = parser.add_argument_group(title="Extra models args, including vae, text encoders and tokenizers)")
# - VAE
group.add_argument(
@@ -98,9 +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,
@@ -195,7 +191,10 @@ def add_denoise_schedule_args(parser: argparse.ArgumentParser):
help="If reverse, learning/sampling from t=1 -> t=0.",
)
group.add_argument(
"--flow-solver", type=str, default="euler", help="Solver for flow matching.",
"--flow-solver",
type=str,
default="euler",
help="Solver for flow matching.",
)
group.add_argument(
"--use-linear-quadratic-schedule",
@@ -330,17 +329,13 @@ def add_inference_args(parser: argparse.ArgumentParser):
group.add_argument("--seed", type=int, default=None, help="Seed for evaluation.")
# Classifier-Free Guidance
group.add_argument(
"--neg-prompt", type=str, default=None, help="Negative prompt for sampling."
)
group.add_argument(
"--cfg-scale", type=float, default=1.0, help="Classifier free guidance scale."
)
group.add_argument("--neg-prompt", type=str, default=None, help="Negative prompt for sampling.")
group.add_argument("--cfg-scale", type=float, default=1.0, help="Classifier free guidance scale.")
group.add_argument(
"--embedded-cfg-scale",
type=float,
default=6.0,
help="Embeded classifier free guidance scale.",
help="Embedded classifier free guidance scale.",
)
group.add_argument(
@@ -357,10 +352,16 @@ def add_parallel_args(parser: argparse.ArgumentParser):
# ======================== Model loads ========================
group.add_argument(
"--ulysses-degree", type=int, default=1, help="Ulysses degree.",
"--ulysses-degree",
type=int,
default=1,
help="Ulysses degree.",
)
group.add_argument(
"--ring-degree", type=int, default=1, help="Ulysses degree.",
"--ring-degree",
type=int,
default=1,
help="Ulysses degree.",
)
return parser
@@ -370,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
+44 -94
View File
@@ -1,33 +1,24 @@
import os
import time
import random
import functools
from typing import List, Optional, Tuple, Union
import time
from pathlib import Path
from loguru import logger
import torch
import torch.distributed as dist
from fastvideo.models.hunyuan.constants import (
PROMPT_TEMPLATE,
NEGATIVE_PROMPT,
PRECISION_TO_TYPE,
)
from fastvideo.models.hunyuan.vae import load_vae
from loguru import logger
from safetensors.torch import load_file as safetensors_load_file
from fastvideo.models.hunyuan.constants import NEGATIVE_PROMPT, PRECISION_TO_TYPE, PROMPT_TEMPLATE
from fastvideo.models.hunyuan.diffusion.pipelines import HunyuanVideoPipeline
from fastvideo.models.hunyuan.diffusion.schedulers import FlowMatchDiscreteScheduler
from fastvideo.models.hunyuan.modules import load_model
from fastvideo.models.hunyuan.text_encoder import TextEncoder
from fastvideo.models.hunyuan.utils.data_utils import align_to
from fastvideo.models.hunyuan.diffusion.schedulers import FlowMatchDiscreteScheduler
from fastvideo.models.hunyuan.diffusion.pipelines import HunyuanVideoPipeline
from safetensors.torch import load_file as safetensors_load_file
from fastvideo.utils.parallel_states import (
initialize_sequence_parallel_state,
nccl_info,
)
from fastvideo.models.hunyuan.vae import load_vae
from fastvideo.utils.parallel_states import nccl_info
class Inference(object):
def __init__(
self,
args,
@@ -53,13 +44,7 @@ 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
@@ -103,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 ========================
@@ -117,9 +104,7 @@ class Inference(object):
# Text encoder
if args.prompt_template_video is not None:
crop_start = PROMPT_TEMPLATE[args.prompt_template_video].get(
"crop_start", 0
)
crop_start = PROMPT_TEMPLATE[args.prompt_template_video].get("crop_start", 0)
elif args.prompt_template is not None:
crop_start = PROMPT_TEMPLATE[args.prompt_template].get("crop_start", 0)
else:
@@ -127,18 +112,11 @@ class Inference(object):
max_length = args.text_len + crop_start
# prompt_template
prompt_template = (
PROMPT_TEMPLATE[args.prompt_template]
if args.prompt_template is not None
else None
)
prompt_template = (PROMPT_TEMPLATE[args.prompt_template] if args.prompt_template is not None else None)
# prompt_template_video
prompt_template_video = (
PROMPT_TEMPLATE[args.prompt_template_video]
if args.prompt_template_video is not None
else None
)
prompt_template_video = (PROMPT_TEMPLATE[args.prompt_template_video]
if args.prompt_template_video is not None else None)
text_encoder = TextEncoder(
text_encoder_type=args.text_encoder,
@@ -195,18 +173,14 @@ class Inference(object):
files = [f for f in files if str(f).endswith("_model_states.pt")]
model_path = files[0]
if len(files) > 1:
logger.warning(
f"Multiple model weights found in {dit_weight}, using {model_path}"
)
logger.warning(f"Multiple model weights found in {dit_weight}, using {model_path}")
bare_model = False
else:
raise ValueError(
f"Invalid model path: {dit_weight} with unrecognized weight format: "
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
f"specific weight file, please provide the full path to the file."
)
raise ValueError(f"Invalid model path: {dit_weight} with unrecognized weight format: "
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
f"specific weight file, please provide the full path to the file.")
else:
if dit_weight.is_dir():
files = list(dit_weight.glob("*.pt"))
@@ -219,18 +193,14 @@ class Inference(object):
files = [f for f in files if str(f).endswith("_model_states.pt")]
model_path = files[0]
if len(files) > 1:
logger.warning(
f"Multiple model weights found in {dit_weight}, using {model_path}"
)
logger.warning(f"Multiple model weights found in {dit_weight}, using {model_path}")
bare_model = False
else:
raise ValueError(
f"Invalid model path: {dit_weight} with unrecognized weight format: "
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
f"specific weight file, please provide the full path to the file."
)
raise ValueError(f"Invalid model path: {dit_weight} with unrecognized weight format: "
f"{list(map(str, files))}. When given a directory as --dit-weight, only "
f"`pytorch_model_*.pt`(provided by HunyuanDiT official) and "
f"`*_model_states.pt`(saved by deepspeed) can be parsed. If you want to load a "
f"specific weight file, please provide the full path to the file.")
elif dit_weight.is_file():
model_path = dit_weight
bare_model = "unknown"
@@ -245,9 +215,7 @@ 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}")
@@ -257,10 +225,8 @@ class Inference(object):
if load_key in state_dict:
state_dict = state_dict[load_key]
else:
raise KeyError(
f"Missing key: `{load_key}` in the checkpoint: {model_path}. The keys in the checkpoint "
f"are: {list(state_dict.keys())}."
)
raise KeyError(f"Missing key: `{load_key}` in the checkpoint: {model_path}. The keys in the checkpoint "
f"are: {list(state_dict.keys())}.")
model.load_state_dict(state_dict, strict=True)
return model
@@ -278,6 +244,7 @@ class Inference(object):
class HunyuanVideoSampler(Inference):
def __init__(
self,
args,
@@ -371,6 +338,7 @@ class HunyuanVideoSampler(Inference):
embedded_guidance_scale=None,
batch_size=1,
num_videos_per_prompt=1,
mask_strategy=None,
**kwargs,
):
"""
@@ -397,34 +365,20 @@ class HunyuanVideoSampler(Inference):
if isinstance(seed, torch.Tensor):
seed = seed.tolist()
if seed is None:
seeds = [
random.randint(0, 1_000_000)
for _ in range(batch_size * num_videos_per_prompt)
]
seeds = [random.randint(0, 1_000_000) for _ in range(batch_size * num_videos_per_prompt)]
elif isinstance(seed, int):
seeds = [
seed + i
for _ in range(batch_size)
for i in range(num_videos_per_prompt)
]
seeds = [seed + i for _ in range(batch_size) for i in range(num_videos_per_prompt)]
elif isinstance(seed, (list, tuple)):
if len(seed) == batch_size:
seeds = [
int(seed[i]) + j
for i in range(batch_size)
for j in range(num_videos_per_prompt)
]
seeds = [int(seed[i]) + j for i in range(batch_size) for j in range(num_videos_per_prompt)]
elif len(seed) == batch_size * num_videos_per_prompt:
seeds = [int(s) for s in seed]
else:
raise ValueError(
f"Length of seed must be equal to number of prompt(batch_size) or "
f"batch_size * num_videos_per_prompt ({batch_size} * {num_videos_per_prompt}), got {seed}."
)
f"batch_size * num_videos_per_prompt ({batch_size} * {num_videos_per_prompt}), got {seed}.")
else:
raise ValueError(
f"Seed must be an integer, a list of integers, or None, got {seed}."
)
raise ValueError(f"Seed must be an integer, a list of integers, or None, got {seed}.")
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
generator = [torch.Generator("cpu").manual_seed(seed) for seed in seeds]
out_dict["seeds"] = seeds
@@ -437,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)
@@ -462,9 +412,7 @@ class HunyuanVideoSampler(Inference):
if negative_prompt is None or negative_prompt == "":
negative_prompt = self.default_negative_prompt
if not isinstance(negative_prompt, str):
raise TypeError(
f"`negative_prompt` must be a string, but got {type(negative_prompt)}"
)
raise TypeError(f"`negative_prompt` must be a string, but got {type(negative_prompt)}")
negative_prompt = [negative_prompt.strip()]
# ========================================================================
@@ -522,6 +470,8 @@ class HunyuanVideoSampler(Inference):
is_progress_bar=True,
vae_ver=self.args.vae,
enable_tiling=self.args.vae_tiling,
enable_vae_sp=self.args.vae_sp,
mask_strategy=mask_strategy,
)[0]
out_dict["samples"] = samples
out_dict["prompts"] = prompt
+1 -1
View File
@@ -1,4 +1,4 @@
from .models import HYVideoDiffusionTransformer, HUNYUAN_VIDEO_CONFIG
from .models import HUNYUAN_VIDEO_CONFIG, HYVideoDiffusionTransformer
def load_model(args, in_channels, out_channels, factor_kwargs):
+69 -29
View File
@@ -1,18 +1,25 @@
import importlib.metadata
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
try:
from st_attn import sliding_tile_attention
except ImportError:
print("Could not load Sliding Tile Attention.")
sliding_tile_attention = None
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
from fastvideo.utils.communications import all_gather, all_to_all_4D
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
from fastvideo.utils.communications import all_gather, all_to_all_4D
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
def attention(
q, k, v, drop_rate=0, attn_mask=None, causal=False,
q,
k,
v,
drop_rate=0,
attn_mask=None,
causal=False,
):
qkv = torch.stack([q, k, v], dim=2)
@@ -20,21 +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)
@@ -43,9 +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)
@@ -55,24 +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 = hidden_states.to(query.dtype)
encoder_hidden_states = encoder_hidden_states.to(query.dtype)
@@ -1,7 +1,7 @@
import math
import torch
import torch.nn as nn
from einops import rearrange, repeat
from ..utils.helpers import to_2tuple
@@ -105,11 +105,8 @@ def timestep_embedding(t, dim, max_period=10000):
.. ref_link: https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py
"""
half = dim // 2
freqs = torch.exp(
-math.log(max_period)
* torch.arange(start=0, end=half, dtype=torch.float32)
/ half
).to(device=t.device)
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) /
half).to(device=t.device)
args = t[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
@@ -140,9 +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),
)
@@ -150,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
+6 -18
View File
@@ -6,8 +6,8 @@ from functools import partial
import torch
import torch.nn as nn
from .modulate_layers import modulate
from ..utils.helpers import to_2tuple
from .modulate_layers import modulate
class MLP(nn.Module):
@@ -34,19 +34,11 @@ class MLP(nn.Module):
drop_probs = to_2tuple(drop)
linear_layer = partial(nn.Conv2d, kernel_size=1) if use_conv else nn.Linear
self.fc1 = linear_layer(
in_channels, hidden_channels, bias=bias[0], **factory_kwargs
)
self.fc1 = linear_layer(in_channels, hidden_channels, bias=bias[0], **factory_kwargs)
self.act = act_layer()
self.drop1 = nn.Dropout(drop_probs[0])
self.norm = (
norm_layer(hidden_channels, **factory_kwargs)
if norm_layer is not None
else nn.Identity()
)
self.fc2 = linear_layer(
hidden_channels, out_features, bias=bias[1], **factory_kwargs
)
self.norm = (norm_layer(hidden_channels, **factory_kwargs) if norm_layer is not None else nn.Identity())
self.fc2 = linear_layer(hidden_channels, out_features, bias=bias[1], **factory_kwargs)
self.drop2 = nn.Dropout(drop_probs[1])
def forward(self, x):
@@ -77,16 +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,
+127 -218
View File
@@ -1,29 +1,27 @@
from typing import Any, List, Tuple, Optional, Union, Dict
from einops import rearrange
from typing import Any, Dict, List, Optional, Tuple, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
from diffusers.models import ModelMixin
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.utils.parallel_states import nccl_info
from .activation_layers import get_activation_layer
from .norm_layers import get_norm_layer
from .embed_layers import TimestepEmbedder, PatchEmbed, TextProjection
from .attenion import parallel_attention
from .embed_layers import PatchEmbed, TextProjection, TimestepEmbedder
from .mlp_layers import MLP, FinalLayer, MLPEmbedder
from .modulate_layers import ModulateDiT, apply_gate, modulate
from .norm_layers import get_norm_layer
from .posemb_layers import apply_rotary_emb
from .mlp_layers import MLP, MLPEmbedder, FinalLayer
from .modulate_layers import ModulateDiT, modulate, apply_gate
from .token_refiner import SingleTokenRefiner
from fastvideo.models.hunyuan.modules.posemb_layers import get_nd_rotary_pos_embed
from fastvideo.utils.parallel_states import nccl_info
class MMDoubleStreamBlock(nn.Module):
"""
A multimodal dit block with seperate modulation for
A multimodal dit block with separate modulation for
text and image/video, see more details (SD3): https://arxiv.org/abs/2403.03206
(Flux.1): https://github.com/black-forest-labs/flux
"""
@@ -54,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)
if qk_norm
else nn.Identity()
)
self.img_attn_k_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.img_attn_proj = nn.Linear(
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
)
self.img_attn_q_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.img_attn_k_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.img_attn_proj = nn.Linear(hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs)
self.img_norm2 = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.img_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
self.img_mlp = MLP(
hidden_size,
mlp_hidden_dim,
@@ -93,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)
if qk_norm
else nn.Identity()
)
self.txt_attn_k_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.txt_attn_proj = nn.Linear(
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
)
self.txt_attn_qkv = nn.Linear(hidden_size, hidden_size * 3, bias=qkv_bias, **factory_kwargs)
self.txt_attn_q_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.txt_attn_k_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.txt_attn_proj = nn.Linear(hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs)
self.txt_norm2 = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.txt_norm2 = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
self.txt_mlp = MLP(
hidden_size,
mlp_hidden_dim,
@@ -139,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,
@@ -159,13 +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,9 +142,7 @@ 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),
@@ -185,20 +150,15 @@ class MMDoubleStreamBlock(nn.Module):
)
img_qq, img_kk = apply_rotary_emb(img_q, img_k, freqs_cis, head_first=False)
assert (
img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
assert (img_qq.shape == img_q.shape and img_kk.shape == img_k.shape
), f"img_kk: {img_qq.shape}, img_q: {img_q.shape}, img_kk: {img_kk.shape}, img_k: {img_k.shape}"
img_q, img_k = img_qq, img_kk
# Prepare txt for attention.
txt_modulated = self.txt_norm1(txt)
txt_modulated = modulate(
txt_modulated, shift=txt_mod1_shift, scale=txt_mod1_scale
)
txt_modulated = modulate(txt_modulated, shift=txt_mod1_shift, scale=txt_mod1_scale)
txt_qkv = self.txt_attn_qkv(txt_modulated)
txt_q, txt_k, txt_v = rearrange(
txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num
)
txt_q, txt_k, txt_v = rearrange(txt_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num)
# Apply QK-Norm if needed.
txt_q = self.txt_attn_q_norm(txt_q).to(txt_v)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_v)
@@ -210,34 +170,26 @@ class MMDoubleStreamBlock(nn.Module):
img_q_len=img_q.shape[1],
img_kv_len=img_k.shape[1],
text_mask=text_mask,
mask_strategy=mask_strategy,
)
# attention computation end
img_attn, txt_attn = attn[:, : img.shape[1]], attn[:, img.shape[1] :]
img_attn, txt_attn = attn[:, :img.shape[1]], attn[:, img.shape[1]:]
# Calculate the img bloks.
# Calculate the img blocks.
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 bloks.
# Calculate the txt blocks.
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
@@ -270,32 +222,20 @@ class MMSingleStreamBlock(nn.Module):
head_dim = hidden_size // heads_num
mlp_hidden_dim = int(hidden_size * mlp_width_ratio)
self.mlp_hidden_dim = mlp_hidden_dim
self.scale = qk_scale or head_dim ** -0.5
self.scale = qk_scale or head_dim**-0.5
# qkv and mlp_in
self.linear1 = nn.Linear(
hidden_size, hidden_size * 3 + mlp_hidden_dim, **factory_kwargs
)
self.linear1 = nn.Linear(hidden_size, hidden_size * 3 + mlp_hidden_dim, **factory_kwargs)
# proj and mlp_out
self.linear2 = nn.Linear(
hidden_size + mlp_hidden_dim, hidden_size, **factory_kwargs
)
self.linear2 = nn.Linear(hidden_size + mlp_hidden_dim, hidden_size, **factory_kwargs)
qk_norm_layer = get_norm_layer(qk_norm_type)
self.q_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.k_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.q_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.k_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.pre_norm = nn.LayerNorm(
hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs
)
self.pre_norm = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, **factory_kwargs)
self.mlp_act = get_activation_layer(mlp_act_type)()
self.modulation = ModulateDiT(
@@ -319,12 +259,11 @@ 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)
@@ -334,19 +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}"
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(
@@ -356,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
@@ -458,21 +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":
@@ -491,61 +425,44 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
**factory_kwargs,
)
else:
raise NotImplementedError(
f"Unsupported text_projection: {self.text_projection}"
)
raise NotImplementedError(f"Unsupported text_projection: {self.text_projection}")
# time modulation
self.time_in = TimestepEmbedder(
self.hidden_size, get_activation_layer("silu"), **factory_kwargs
)
self.time_in = TimestepEmbedder(self.hidden_size, get_activation_layer("silu"), **factory_kwargs)
# text modulation
self.vector_in = MLPEmbedder(
self.config.text_states_dim_2, self.hidden_size, **factory_kwargs
)
self.vector_in = MLPEmbedder(self.config.text_states_dim_2, self.hidden_size, **factory_kwargs)
# guidance modulation
self.guidance_in = (
TimestepEmbedder(
self.hidden_size, get_activation_layer("silu"), **factory_kwargs
)
if guidance_embed
else None
)
self.guidance_in = (TimestepEmbedder(self.hidden_size, get_activation_layer("silu"), **factory_kwargs)
if guidance_embed else None)
# double blocks
self.double_blocks = nn.ModuleList(
[
MMDoubleStreamBlock(
self.hidden_size,
self.heads_num,
mlp_width_ratio=mlp_width_ratio,
mlp_act_type=mlp_act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
qkv_bias=qkv_bias,
**factory_kwargs,
)
for _ in range(mm_double_blocks_depth)
]
)
self.double_blocks = nn.ModuleList([
MMDoubleStreamBlock(
self.hidden_size,
self.heads_num,
mlp_width_ratio=mlp_width_ratio,
mlp_act_type=mlp_act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
qkv_bias=qkv_bias,
**factory_kwargs,
) for _ in range(mm_double_blocks_depth)
])
# single blocks
self.single_blocks = nn.ModuleList(
[
MMSingleStreamBlock(
self.hidden_size,
self.heads_num,
mlp_width_ratio=mlp_width_ratio,
mlp_act_type=mlp_act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
**factory_kwargs,
)
for _ in range(mm_single_blocks_depth)
]
)
self.single_blocks = nn.ModuleList([
MMSingleStreamBlock(
self.hidden_size,
self.heads_num,
mlp_width_ratio=mlp_width_ratio,
mlp_act_type=mlp_act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
**factory_kwargs,
) for _ in range(mm_single_blocks_depth)
])
self.final_layer = FinalLayer(
self.hidden_size,
@@ -569,14 +486,12 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
def get_rotary_pos_embed(self, rope_sizes):
target_ndim = 3
ndim = 5 - 2
head_dim = self.hidden_size // self.heads_num
rope_dim_list = self.rope_dim_list
if rope_dim_list is None:
rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)]
assert (
sum(rope_dim_list) == head_dim
), "sum(rope_dim_list) should equal to head_dim of attention layer"
assert (sum(rope_dim_list) == head_dim), "sum(rope_dim_list) should equal to head_dim of attention layer"
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
rope_dim_list,
rope_sizes,
@@ -599,27 +514,27 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
encoder_hidden_states: torch.Tensor,
timestep: torch.LongTensor,
encoder_attention_mask: torch.Tensor,
mask_strategy=None,
output_features=False,
output_features_stride=8,
attention_kwargs: Optional[Dict[str, Any]] = None,
return_dict: bool = False,
guidance=None,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
if guidance == None:
guidance = torch.tensor(
[6016.0], device=hidden_states.device, dtype=torch.bfloat16
)
out = {}
if guidance is None:
guidance = torch.tensor([6016.0], device=hidden_states.device, dtype=torch.bfloat16)
if mask_strategy is None:
mask_strategy = [[None] * self.heads_num for _ in range(len(self.double_blocks) + len(self.single_blocks))]
img = x = hidden_states
text_mask = encoder_attention_mask
t = timestep
txt = encoder_hidden_states[:, 1:]
text_states_2 = encoder_hidden_states[:, 0, : self.config.text_states_dim_2]
_, _, ot, oh, ow = x.shape
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],
oh // self.patch_size[1],
ow // self.patch_size[2],
ot // self.patch_size[0], # codespell:ignore
oh // self.patch_size[1], # codespell:ignore
ow // self.patch_size[2], # codespell:ignore
)
original_tt = nccl_info.sp_size * tt
freqs_cos, freqs_sin = self.get_rotary_pos_embed((original_tt, th, tw))
@@ -632,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)
@@ -646,34 +559,31 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
elif self.text_projection == "single_refiner":
txt = self.txt_in(txt, t, text_mask if self.use_attention_mask else None)
else:
raise NotImplementedError(
f"Unsupported text_projection: {self.text_projection}"
)
raise NotImplementedError(f"Unsupported text_projection: {self.text_projection}")
txt_seq_len = txt.shape[1]
img_seq_len = img.shape[1]
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# --------------------- Pass through DiT blocks ------------------------
for _, block in enumerate(self.double_blocks):
double_block_args = [img, txt, vec, freqs_cis, text_mask]
for index, block in enumerate(self.double_blocks):
double_block_args = [img, txt, vec, freqs_cis, text_mask, mask_strategy[index]]
img, txt = block(*double_block_args)
# Merge txt and img to pass through single stream blocks.
x = torch.cat((img, txt), 1)
if output_features:
features_list = []
if len(self.single_blocks) > 0:
for _, 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, ...])
@@ -684,7 +594,7 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
img = self.unpatchify(img, tt, th, tw)
assert return_dict == False, "return_dict is not supported."
assert not return_dict, "return_dict is not supported."
if output_features:
features_list = torch.stack(features_list, dim=0)
else:
@@ -708,25 +618,24 @@ class HYVideoDiffusionTransformer(ModelMixin, ConfigMixin):
def params_count(self):
counts = {
"double": sum(
[
sum(p.numel() for p in block.img_attn_qkv.parameters())
+ sum(p.numel() for p in block.img_attn_proj.parameters())
+ sum(p.numel() for p in block.img_mlp.parameters())
+ sum(p.numel() for p in block.txt_attn_qkv.parameters())
+ sum(p.numel() for p in block.txt_attn_proj.parameters())
+ sum(p.numel() for p in block.txt_mlp.parameters())
for block in self.double_blocks
]
),
"single": sum(
[
sum(p.numel() for p in block.linear1.parameters())
+ sum(p.numel() for p in block.linear2.parameters())
for block in self.single_blocks
]
),
"total": sum(p.numel() for p in self.parameters()),
"double":
sum([
sum(p.numel()
for p in block.img_attn_qkv.parameters()) + sum(p.numel()
for p in block.img_attn_proj.parameters()) +
sum(p.numel() for p in block.img_mlp.parameters()) + sum(p.numel()
for p in block.txt_attn_qkv.parameters()) +
sum(p.numel() for p in block.txt_attn_proj.parameters()) + sum(p.numel()
for p in block.txt_mlp.parameters())
for block in self.double_blocks
]),
"single":
sum([
sum(p.numel() for p in block.linear1.parameters()) + sum(p.numel() for p in block.linear2.parameters())
for block in self.single_blocks
]),
"total":
sum(p.numel() for p in self.parameters()),
}
counts["attn+mlp"] = counts["double"] + counts["single"]
return counts
@@ -18,9 +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)
@@ -70,6 +68,7 @@ def apply_gate(x, gate=None, tanh=False):
def ckpt_wrapper(module):
def ckpt_forward(*inputs):
outputs = module(*inputs)
return outputs
@@ -77,11 +76,8 @@ def ckpt_wrapper(module):
return ckpt_forward
import torch
import torch.nn as nn
class RMSNorm(nn.Module):
def __init__(
self,
dim: int,
@@ -3,6 +3,7 @@ import torch.nn as nn
class RMSNorm(nn.Module):
def __init__(
self,
dim: int,
@@ -1,10 +1,11 @@
from typing import List, Tuple, Union
import torch
from typing import Union, Tuple, List
def _to_tuple(x, dim=2):
if isinstance(x, int):
return (x,) * dim
return (x, ) * dim
elif len(x) == dim:
return x
else:
@@ -29,7 +30,7 @@ def get_meshgrid_nd(start, *args, dim=2):
if len(args) == 0:
# start is grid_size
num = _to_tuple(start, dim=dim)
start = (0,) * dim
start = (0, ) * dim
stop = num
elif len(args) == 1:
# start is start, args[0] is stop, step is 1
@@ -99,10 +100,7 @@ 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],
@@ -117,10 +115,7 @@ 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],
@@ -131,9 +126,7 @@ def reshape_for_broadcast(
def rotate_half(x):
x_real, x_imag = (
x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1)
) # [B, S, H, D//2]
x_real, x_imag = (x.float().reshape(*x.shape[:-1], -1, 2).unbind(-1)) # [B, S, H, D//2]
return torch.stack([-x_imag, x_real], dim=-1).flatten(3)
@@ -171,18 +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
@@ -216,25 +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):
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
elif isinstance(theta_rescale_factor, list) and len(theta_rescale_factor) == 1:
theta_rescale_factor = [theta_rescale_factor[0]] * len(rope_dim_list)
assert len(theta_rescale_factor) == len(
rope_dim_list
), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
rope_dim_list), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
if isinstance(interpolation_factor, int) or isinstance(interpolation_factor, float):
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
elif isinstance(interpolation_factor, list) and len(interpolation_factor) == 1:
interpolation_factor = [interpolation_factor[0]] * len(rope_dim_list)
assert len(interpolation_factor) == len(
rope_dim_list
), "len(interpolation_factor) should equal to len(rope_dim_list)"
rope_dim_list), "len(interpolation_factor) should equal to len(rope_dim_list)"
# use 1/ndim of dimensions to encode grid_axis
embs = []
@@ -292,11 +275,9 @@ def get_1d_rotary_pos_embed(
# proposed by reddit user bloc97, to rescale rotary embeddings to longer sequence length without fine-tuning
# has some connection to NTK literature
if theta_rescale_factor != 1.0:
theta *= theta_rescale_factor ** (dim / (dim - 2))
theta *= theta_rescale_factor**(dim / (dim - 2))
freqs = 1.0 / (
theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)
) # [D/2]
freqs = 1.0 / (theta**(torch.arange(0, dim, 2)[:(dim // 2)].float() / dim)) # [D/2]
# assert interpolation_factor == 1.0, f"interpolation_factor: {interpolation_factor}"
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
if use_real:
@@ -304,7 +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
@@ -1,19 +1,19 @@
from typing import Optional
from einops import rearrange
import torch
import torch.nn as nn
from einops import rearrange
from .activation_layers import get_activation_layer
from .attenion import attention
from .norm_layers import get_norm_layer
from .embed_layers import TimestepEmbedder, TextProjection
from .attenion import attention
from .embed_layers import TextProjection, TimestepEmbedder
from .mlp_layers import MLP
from .modulate_layers import modulate, apply_gate
from .modulate_layers import apply_gate
from .norm_layers import get_norm_layer
class IndividualTokenRefinerBlock(nn.Module):
def __init__(
self,
hidden_size,
@@ -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)
if qk_norm
else nn.Identity()
)
self.self_attn_k_norm = (
qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm
else nn.Identity()
)
self.self_attn_proj = nn.Linear(
hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs
)
self.self_attn_q_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.self_attn_k_norm = (qk_norm_layer(head_dim, elementwise_affine=True, eps=1e-6, **factory_kwargs)
if qk_norm else nn.Identity())
self.self_attn_proj = nn.Linear(hidden_size, hidden_size, bias=qkv_bias, **factory_kwargs)
self.norm2 = nn.LayerNorm(
hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs
)
self.norm2 = nn.LayerNorm(hidden_size, elementwise_affine=True, eps=1e-6, **factory_kwargs)
act_layer = get_activation_layer(act_type)
self.mlp = MLP(
in_channels=hidden_size,
@@ -101,6 +87,7 @@ class IndividualTokenRefinerBlock(nn.Module):
class IndividualTokenRefiner(nn.Module):
def __init__(
self,
hidden_size,
@@ -117,25 +104,25 @@ class IndividualTokenRefiner(nn.Module):
):
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.blocks = nn.ModuleList(
[
IndividualTokenRefinerBlock(
hidden_size=hidden_size,
heads_num=heads_num,
mlp_width_ratio=mlp_width_ratio,
mlp_drop_rate=mlp_drop_rate,
act_type=act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
qkv_bias=qkv_bias,
**factory_kwargs,
)
for _ in range(depth)
]
)
self.blocks = nn.ModuleList([
IndividualTokenRefinerBlock(
hidden_size=hidden_size,
heads_num=heads_num,
mlp_width_ratio=mlp_width_ratio,
mlp_drop_rate=mlp_drop_rate,
act_type=act_type,
qk_norm=qk_norm,
qk_norm_type=qk_norm_type,
qkv_bias=qkv_bias,
**factory_kwargs,
) for _ in range(depth)
])
def forward(
self, x: torch.Tensor, c: torch.LongTensor, mask: Optional[torch.Tensor] = None,
self,
x: torch.Tensor,
c: torch.LongTensor,
mask: Optional[torch.Tensor] = None,
):
mask = mask.clone().bool()
# avoid attention weight become NaN
@@ -171,17 +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)
# 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,
@@ -208,9 +191,7 @@ class SingleTokenRefiner(nn.Module):
context_aware_representations = x.mean(dim=1)
else:
mask_float = mask.float().unsqueeze(-1) # [b, s1, 1]
context_aware_representations = (x * mask_float).sum(
dim=1
) / mask_float.sum(dim=1)
context_aware_representations = (x * mask_float).sum(dim=1) / mask_float.sum(dim=1)
context_aware_representations = self.c_embedder(context_aware_representations)
c = timestep_aware_representations + context_aware_representations
@@ -16,7 +16,6 @@ Given Input:
input: "{input}"
"""
master_mode_prompt = """Master mode - Video Recaption Task:
You are a large language model specialized in rewriting video descriptions. Your task is to modify the input description.
@@ -1,14 +1,12 @@
from dataclasses import dataclass
from typing import Optional, Tuple
from copy import deepcopy
import torch
import torch.nn as nn
from transformers import CLIPTextModel, CLIPTokenizer, AutoTokenizer, AutoModel
from transformers import AutoModel, AutoTokenizer, CLIPTextModel, CLIPTokenizer
from transformers.utils import ModelOutput
from ..constants import TEXT_ENCODER_PATH, TOKENIZER_PATH
from ..constants import PRECISION_TO_TYPE
from ..constants import PRECISION_TO_TYPE, TEXT_ENCODER_PATH, TOKENIZER_PATH
def use_default(value, default):
@@ -25,17 +23,13 @@ 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}")
@@ -55,9 +49,7 @@ 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:
@@ -66,9 +58,7 @@ def load_tokenizer(
if tokenizer_type == "clipL":
tokenizer = CLIPTokenizer.from_pretrained(tokenizer_path, max_length=77)
elif tokenizer_type == "llm":
tokenizer = AutoTokenizer.from_pretrained(
tokenizer_path, padding_side=padding_side
)
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path, padding_side=padding_side)
else:
raise ValueError(f"Unsupported tokenizer type: {tokenizer_type}")
@@ -100,6 +90,7 @@ class TextEncoderModelOutput(ModelOutput):
class TextEncoder(nn.Module):
def __init__(
self,
text_encoder_type: str,
@@ -124,20 +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
@@ -147,26 +130,21 @@ class TextEncoder(nn.Module):
self.use_template = self.prompt_template is not None
if self.use_template:
assert (
isinstance(self.prompt_template, dict)
and "template" in self.prompt_template
), f"`prompt_template` must be a dictionary with a key 'template', got {self.prompt_template}"
assert (isinstance(self.prompt_template, dict) and "template" in self.prompt_template
), f"`prompt_template` must be a dictionary with a key 'template', got {self.prompt_template}"
assert "{}" in str(self.prompt_template["template"]), (
"`prompt_template['template']` must contain a placeholder `{}` for the input text, "
f"got {self.prompt_template['template']}"
)
f"got {self.prompt_template['template']}")
self.use_video_template = self.prompt_template_video is not None
if self.use_video_template:
if self.prompt_template_video is not None:
assert (
isinstance(self.prompt_template_video, dict)
and "template" in self.prompt_template_video
isinstance(self.prompt_template_video, dict) and "template" in self.prompt_template_video
), f"`prompt_template_video` must be a dictionary with a key 'template', got {self.prompt_template_video}"
assert "{}" in str(self.prompt_template_video["template"]), (
"`prompt_template_video['template']` must contain a placeholder `{}` for the input text, "
f"got {self.prompt_template_video['template']}"
)
f"got {self.prompt_template_video['template']}")
if "t5" in text_encoder_type:
self.output_key = output_key or "last_hidden_state"
@@ -205,7 +183,7 @@ class TextEncoder(nn.Module):
Args:
text (str): Input text.
template (str or list): Template string or list of chat conversation.
prevent_empty_text (bool): If Ture, we will prevent the user text from being empty
prevent_empty_text (bool): If True, we will prevent the user text from being empty
by adding a space. Defaults to True.
"""
if isinstance(template, str):
@@ -230,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):
@@ -295,18 +270,13 @@ class TextEncoder(nn.Module):
"""
device = self.model.device if device is None else device
use_attention_mask = use_default(use_attention_mask, self.use_attention_mask)
hidden_state_skip_layer = use_default(
hidden_state_skip_layer, self.hidden_state_skip_layer
)
hidden_state_skip_layer = use_default(hidden_state_skip_layer, self.hidden_state_skip_layer)
do_sample = use_default(do_sample, not self.reproduce)
attention_mask = (
batch_encoding["attention_mask"].to(device) if use_attention_mask else None
)
attention_mask = (batch_encoding["attention_mask"].to(device) if use_attention_mask else None)
outputs = self.model(
input_ids=batch_encoding["input_ids"].to(device),
attention_mask=attention_mask,
output_hidden_states=output_hidden_states
or hidden_state_skip_layer is not None,
output_hidden_states=output_hidden_states or hidden_state_skip_layer is not None,
)
if hidden_state_skip_layer is not None:
last_hidden_state = outputs.hidden_states[-(hidden_state_skip_layer + 1)]
@@ -327,14 +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(
+1 -2
View File
@@ -1,9 +1,8 @@
import numpy as np
import math
def align_to(value, alignment):
"""align hight, width according to alignment
"""align height, width according to alignment
Args:
value (int): height or width
+3 -3
View File
@@ -1,11 +1,11 @@
import os
from pathlib import Path
from einops import rearrange
import imageio
import numpy as np
import torch
import torchvision
import numpy as np
import imageio
from einops import rearrange
CODE_SUFFIXES = {
".py", # Python codes
+2 -2
View File
@@ -1,9 +1,9 @@
import collections.abc
from itertools import repeat
def _ntuple(n):
def parse(x):
if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):
x = tuple(x)
@@ -25,7 +25,7 @@ def as_tuple(x):
if isinstance(x, collections.abc.Iterable) and not isinstance(x, str):
return tuple(x)
if x is None or isinstance(x, (int, float, str)):
return (x,)
return (x, )
else:
raise ValueError(f"Unknown type {type(x)}")
@@ -1,16 +1,16 @@
import argparse
import torch
from transformers import (
AutoProcessor,
LlavaForConditionalGeneration,
)
from transformers import AutoProcessor, LlavaForConditionalGeneration
def preprocess_text_encoder_tokenizer(args):
processor = AutoProcessor.from_pretrained(args.input_dir)
model = LlavaForConditionalGeneration.from_pretrained(
args.input_dir, torch_dtype=torch.float16, low_cpu_mem_usage=True,
args.input_dir,
torch_dtype=torch.float16,
low_cpu_mem_usage=True,
).to(0)
model.language_model.save_pretrained(f"{args.output_dir}")

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