Compare commits
34
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c83432bdea | ||
|
|
e90d7595b1 | ||
|
|
83db422e9c | ||
|
|
dc2a4514e8 | ||
|
|
dd10588fb7 | ||
|
|
8140cb269f | ||
|
|
bb6ae368f2 | ||
|
|
e14d384a71 | ||
|
|
6c528468b5 | ||
|
|
ae5ed0c7e0 | ||
|
|
c37535ab0b | ||
|
|
4e056b92cd | ||
|
|
83348c6ed0 | ||
|
|
691f9d1064 | ||
|
|
796eaf809f | ||
|
|
f51e9d486b | ||
|
|
759f243cf3 | ||
|
|
fb44fbaa1c | ||
|
|
e976583cf4 | ||
|
|
d2db0d475b | ||
|
|
d5ac1e9bee | ||
|
|
1976b23121 | ||
|
|
eafeea4a3f | ||
|
|
bf1fc27989 | ||
|
|
8d99ec3c85 | ||
|
|
b631546e18 | ||
|
|
b2def4b57c | ||
|
|
23fd3ed3c7 | ||
|
|
6e8b11c137 | ||
|
|
fc5a4bc236 | ||
|
|
0f4c8d1360 | ||
|
|
5252d50b25 | ||
|
|
ac07e436bb | ||
|
|
42f902cf23 |
@@ -1,240 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
|
||||
import requests
|
||||
|
||||
|
||||
def parse_arguments():
|
||||
"""Parse command line arguments"""
|
||||
parser = argparse.ArgumentParser(description='Run tests on RunPod GPU')
|
||||
parser.add_argument('--gpu-type', type=str, help='GPU type to use')
|
||||
parser.add_argument('--gpu-count',
|
||||
type=int,
|
||||
help='Number of GPUs to use',
|
||||
default=1)
|
||||
parser.add_argument('--test-command', type=str, help='Test command to run')
|
||||
parser.add_argument('--disk-size',
|
||||
type=int,
|
||||
default=20,
|
||||
help='Container disk size in GB (default: 20)')
|
||||
parser.add_argument('--volume-size',
|
||||
type=int,
|
||||
default=20,
|
||||
help='Persistent volume size in GB (default: 20)')
|
||||
parser.add_argument(
|
||||
'--image',
|
||||
type=str,
|
||||
default='runpod/pytorch:2.4.0-py3.11-cuda12.4.1-devel-ubuntu22.04',
|
||||
help='Docker image to use')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
args = parse_arguments()
|
||||
API_KEY = os.environ['RUNPOD_API_KEY']
|
||||
RUN_ID = os.environ['GITHUB_RUN_ID']
|
||||
JOB_ID = os.environ['JOB_ID']
|
||||
PODS_API = "https://rest.runpod.io/v1/pods"
|
||||
HEADERS = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {API_KEY}"
|
||||
}
|
||||
|
||||
|
||||
def create_pod():
|
||||
"""Create a RunPod instance"""
|
||||
print(f"Creating RunPod instance with GPU: {args.gpu_type}...")
|
||||
payload = {
|
||||
"name": f"fastvideo-{JOB_ID}-{RUN_ID}",
|
||||
"containerDiskInGb": args.disk_size,
|
||||
"volumeInGb": args.volume_size,
|
||||
"gpuTypeIds": [args.gpu_type],
|
||||
"gpuCount": args.gpu_count,
|
||||
"imageName": args.image,
|
||||
"allowedCudaVersions": ["12.4"]
|
||||
}
|
||||
|
||||
response = requests.post(PODS_API, headers=HEADERS, json=payload)
|
||||
response_data = response.json()
|
||||
print(f"Response: {json.dumps(response_data, indent=2)}")
|
||||
|
||||
return response_data["id"]
|
||||
|
||||
|
||||
def wait_for_pod(pod_id):
|
||||
"""Wait for pod to be in RUNNING state and fully ready with SSH access"""
|
||||
print("Waiting for RunPod to be ready...")
|
||||
|
||||
# First wait for RUNNING status
|
||||
max_attempts = 10
|
||||
attempts = 0
|
||||
while attempts < max_attempts:
|
||||
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
pod_data = response.json()
|
||||
status = pod_data["desiredStatus"]
|
||||
|
||||
if status == "RUNNING":
|
||||
print("RunPod is running! Now waiting for ports to be assigned...")
|
||||
break
|
||||
|
||||
print(
|
||||
f"Current status: {status}, waiting... (attempt {attempts+1}/{max_attempts})"
|
||||
)
|
||||
time.sleep(2)
|
||||
attempts += 1
|
||||
|
||||
if attempts >= max_attempts:
|
||||
raise TimeoutError(
|
||||
"Timed out waiting for RunPod to reach RUNNING state")
|
||||
|
||||
# Wait for ports to be assigned
|
||||
max_attempts = 6
|
||||
attempts = 0
|
||||
while attempts < max_attempts:
|
||||
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
pod_data = response.json()
|
||||
port_mappings = pod_data.get("portMappings")
|
||||
|
||||
if (port_mappings is not None and "22" in port_mappings
|
||||
and pod_data.get("publicIp", "") != ""):
|
||||
print("RunPod is ready with SSH access!")
|
||||
print(f"SSH IP: {pod_data['publicIp']}")
|
||||
print(f"SSH Port: {port_mappings['22']}")
|
||||
break
|
||||
|
||||
print(
|
||||
f"Waiting for SSH port and public IP to be available... (attempt {attempts+1}/{max_attempts})"
|
||||
)
|
||||
time.sleep(10)
|
||||
attempts += 1
|
||||
|
||||
if attempts >= max_attempts:
|
||||
raise TimeoutError("Timed out waiting for RunPod SSH access")
|
||||
|
||||
|
||||
def execute_command(pod_id):
|
||||
"""Execute command on the pod via SSH using system SSH client"""
|
||||
print(f"Running command: {args.test_command}")
|
||||
|
||||
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
pod_data = response.json()
|
||||
ssh_ip = pod_data["publicIp"]
|
||||
ssh_port = pod_data["portMappings"]["22"]
|
||||
|
||||
# Copy the repository to the pod using scp
|
||||
repo_dir = os.path.abspath(os.getcwd())
|
||||
repo_name = os.path.basename(repo_dir)
|
||||
|
||||
print(f"Copying repository from {repo_dir} to RunPod...")
|
||||
|
||||
tar_command = [
|
||||
"tar", "-czf", "/tmp/repo.tar.gz", "-C",
|
||||
os.path.dirname(repo_dir), repo_name
|
||||
]
|
||||
subprocess.run(tar_command, check=True)
|
||||
|
||||
# Copy the tarball to the pod
|
||||
scp_command = [
|
||||
"scp", "-o", "StrictHostKeyChecking=no", "-o",
|
||||
"UserKnownHostsFile=/dev/null", "-o", "ServerAliveInterval=60", "-o",
|
||||
"ServerAliveCountMax=10", "-P",
|
||||
str(ssh_port), "/tmp/repo.tar.gz", f"root@{ssh_ip}:/tmp/"
|
||||
]
|
||||
subprocess.run(scp_command, check=True)
|
||||
|
||||
setup_steps = [
|
||||
"cd /workspace",
|
||||
"wget -q https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh",
|
||||
"bash Miniconda3-latest-Linux-x86_64.sh -b -p $HOME/miniconda3",
|
||||
"source $HOME/miniconda3/bin/activate",
|
||||
"conda create --name venv python=3.10.0 -y", "conda activate venv",
|
||||
"mkdir -p /workspace/repo",
|
||||
"tar -xzf /tmp/repo.tar.gz --no-same-owner -C /workspace/",
|
||||
f"cd /workspace/{repo_name}", args.test_command
|
||||
]
|
||||
remote_command = " && ".join(setup_steps)
|
||||
|
||||
ssh_command = [
|
||||
"ssh", "-o", "StrictHostKeyChecking=no", "-o",
|
||||
"UserKnownHostsFile=/dev/null", "-o", "ServerAliveInterval=60", "-o",
|
||||
"ServerAliveCountMax=10", "-p",
|
||||
str(ssh_port), f"root@{ssh_ip}", remote_command
|
||||
]
|
||||
|
||||
print(f"Connecting to {ssh_ip}:{ssh_port}...")
|
||||
|
||||
try:
|
||||
process = subprocess.Popen(ssh_command,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=0)
|
||||
|
||||
stdout_lines = []
|
||||
|
||||
print("Command output:")
|
||||
|
||||
for line in iter(process.stdout.readline, ''):
|
||||
print(line.strip())
|
||||
stdout_lines.append(line)
|
||||
|
||||
process.wait()
|
||||
|
||||
return_code = process.returncode
|
||||
success = return_code == 0
|
||||
|
||||
stdout_str = "".join(stdout_lines)
|
||||
|
||||
if success:
|
||||
print("Command executed successfully")
|
||||
else:
|
||||
print(f"Command failed with exit code {return_code}")
|
||||
|
||||
result = {
|
||||
"success": success,
|
||||
"return_code": return_code,
|
||||
"stdout": stdout_str,
|
||||
"stderr": ""
|
||||
}
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error executing SSH command: {str(e)}")
|
||||
result = {"success": False, "error": str(e), "stdout": "", "stderr": ""}
|
||||
return result
|
||||
|
||||
|
||||
def terminate_pod(pod_id):
|
||||
"""Terminate the pod"""
|
||||
print("Terminating RunPod...")
|
||||
requests.delete(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
print(f"Terminated pod {pod_id}")
|
||||
|
||||
|
||||
def main():
|
||||
pod_id = None
|
||||
try:
|
||||
pod_id = create_pod()
|
||||
wait_for_pod(pod_id)
|
||||
result = execute_command(pod_id)
|
||||
|
||||
if result.get("error") is not None:
|
||||
print(f"Error executing command: {result['error']}")
|
||||
sys.exit(1)
|
||||
|
||||
if not result.get("success", False):
|
||||
print(
|
||||
"Tests failed - check the output above for details on which tests failed"
|
||||
)
|
||||
sys.exit(1)
|
||||
|
||||
finally:
|
||||
if pod_id:
|
||||
terminate_pod(pod_id)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,90 +0,0 @@
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import uuid
|
||||
|
||||
import requests
|
||||
|
||||
API_KEY = os.environ['RUNPOD_API_KEY']
|
||||
RUN_ID = os.environ.get('GITHUB_RUN_ID', str(uuid.uuid4()))
|
||||
PODS_API = "https://rest.runpod.io/v1/pods"
|
||||
HEADERS = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {API_KEY}"
|
||||
}
|
||||
|
||||
|
||||
def get_job_ids():
|
||||
"""Parse job IDs from environment variable"""
|
||||
job_ids_str = os.environ.get('JOB_IDS')
|
||||
try:
|
||||
job_ids = json.loads(job_ids_str)
|
||||
if not isinstance(job_ids, list):
|
||||
print("Error: JOB_IDS is not a list.")
|
||||
sys.exit(1)
|
||||
return job_ids
|
||||
except json.JSONDecodeError as e:
|
||||
print(f"Error parsing JOB_IDS: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def cleanup_pods():
|
||||
"""Find and terminate RunPod instances"""
|
||||
print(f"Run ID: {RUN_ID}")
|
||||
|
||||
single_job_id = os.environ.get('JOB_ID')
|
||||
|
||||
if single_job_id:
|
||||
job_ids = [single_job_id]
|
||||
print(f"Job ID: {single_job_id}")
|
||||
else:
|
||||
job_ids = get_job_ids()
|
||||
print(f"Job IDs: {job_ids}")
|
||||
|
||||
# Get all pods associated with RunPod API_KEY
|
||||
try:
|
||||
response = requests.get(PODS_API, headers=HEADERS)
|
||||
response.raise_for_status()
|
||||
pods = response.json()
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"Error getting pods: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
# Find and terminate pods created by this workflow run
|
||||
terminated_pods = []
|
||||
for pod in pods:
|
||||
pod_name = pod.get("name", "")
|
||||
pod_id = pod.get("id")
|
||||
|
||||
# Check if this pod was created by one of our jobs
|
||||
if any(f"{job_id}-{RUN_ID}" in pod_name for job_id in job_ids):
|
||||
print(f"Found pod: {pod_id} ({pod_name})")
|
||||
try:
|
||||
print(f"Terminating pod {pod_id}...")
|
||||
term_response = requests.delete(f"{PODS_API}/{pod_id}",
|
||||
headers=HEADERS)
|
||||
term_response.raise_for_status()
|
||||
terminated_pods.append(pod_id)
|
||||
print(f"Successfully terminated pod {pod_id}")
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"Error terminating pod {pod_id}: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
if terminated_pods:
|
||||
if single_job_id:
|
||||
print(f"Terminated pod: {terminated_pods[0]}")
|
||||
else:
|
||||
print(f"Terminated {len(terminated_pods)} pods: {terminated_pods}")
|
||||
else:
|
||||
if single_job_id:
|
||||
print(f"No pod found matching pattern: {single_job_id}-{RUN_ID}")
|
||||
else:
|
||||
print("No pods found to terminate.")
|
||||
|
||||
|
||||
def main():
|
||||
cleanup_pods()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,45 @@
|
||||
name: codespell
|
||||
|
||||
on:
|
||||
# Trigger the workflow on push or pull request,
|
||||
# but only for the main branch
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- "**/*.md"
|
||||
- "**/*.rst"
|
||||
- pyproject.toml
|
||||
- requirements-lint.txt
|
||||
- .github/workflows/codespell.yml
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- "**/*.md"
|
||||
- "**/*.rst"
|
||||
- pyproject.toml
|
||||
- requirements-lint.txt
|
||||
- .github/workflows/codespell.yml
|
||||
|
||||
jobs:
|
||||
codespell:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v3
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: '3.12' # or any version you need
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements-lint.txt
|
||||
- name: Spelling check with codespell
|
||||
run: |
|
||||
# Refer to the above environment variable here
|
||||
codespell --toml pyproject.toml $CODESPELL_EXCLUDES
|
||||
@@ -1,77 +0,0 @@
|
||||
# Sample workflow for building and deploying a Hugo site to GitHub Pages
|
||||
name: Deploy FastVideo Docs to Pages
|
||||
|
||||
on:
|
||||
# Runs on pushes targeting the default branch
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "docs/**/*.md"
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
types: [opened, ready_for_review, synchronize, reopened]
|
||||
paths:
|
||||
- "docs/**/*.md"
|
||||
|
||||
# Allows you to run this workflow manually from the Actions tab
|
||||
workflow_dispatch:
|
||||
|
||||
# Sets permissions of the GITHUB_TOKEN to allow deployment to GitHub Pages
|
||||
permissions:
|
||||
contents: read
|
||||
pages: write
|
||||
id-token: write
|
||||
|
||||
# Allow only one concurrent deployment, skipping runs queued between the run in-progress and latest queued.
|
||||
# However, do NOT cancel in-progress runs as we want to allow these production deployments to complete.
|
||||
concurrency:
|
||||
group: "pages"
|
||||
cancel-in-progress: false
|
||||
|
||||
# Default to bash
|
||||
defaults:
|
||||
run:
|
||||
shell: bash
|
||||
|
||||
jobs:
|
||||
# Build job
|
||||
build:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
- name: Setup Pages
|
||||
id: pages
|
||||
uses: actions/configure-pages@v5
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
cd docs
|
||||
pip install -r requirements-docs.txt
|
||||
- name: Build docs
|
||||
run: |
|
||||
cd docs
|
||||
make clean
|
||||
make html
|
||||
- name: Upload artifact
|
||||
uses: actions/upload-pages-artifact@v3
|
||||
with:
|
||||
path: ./docs/build/html
|
||||
|
||||
# Deployment job
|
||||
deploy:
|
||||
environment:
|
||||
name: github-pages
|
||||
url: ${{ steps.deployment.outputs.page_url }}
|
||||
if: ${{ github.event_name == 'push' }}
|
||||
runs-on: ubuntu-latest
|
||||
needs: build
|
||||
steps:
|
||||
- name: Deploy to GitHub Pages
|
||||
id: deployment
|
||||
uses: actions/deploy-pages@v4
|
||||
@@ -15,7 +15,7 @@ jobs:
|
||||
new-version: ${{ steps.check-version.outputs.new-version }}
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
uses: actions/checkout@v3
|
||||
with:
|
||||
fetch-depth: 2
|
||||
|
||||
@@ -23,11 +23,11 @@ jobs:
|
||||
id: check-version
|
||||
run: |
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP "version\\s*=\\s*\"\\K[^\"]+\"" pyproject.toml)
|
||||
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")
|
||||
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
|
||||
@@ -48,10 +48,10 @@ jobs:
|
||||
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
uses: actions/checkout@v3
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
|
||||
@@ -1,17 +0,0 @@
|
||||
{
|
||||
"problemMatcher": [
|
||||
{
|
||||
"owner": "actionlint",
|
||||
"pattern": [
|
||||
{
|
||||
"regexp": "^(?:\\x1b\\[\\d+m)?(.+?)(?:\\x1b\\[\\d+m)*:(?:\\x1b\\[\\d+m)*(\\d+)(?:\\x1b\\[\\d+m)*:(?:\\x1b\\[\\d+m)*(\\d+)(?:\\x1b\\[\\d+m)*: (?:\\x1b\\[\\d+m)*(.+?)(?:\\x1b\\[\\d+m)* \\[(.+?)\\]$",
|
||||
"file": 1,
|
||||
"line": 2,
|
||||
"column": 3,
|
||||
"message": 4,
|
||||
"code": 5
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -1,16 +0,0 @@
|
||||
{
|
||||
"problemMatcher": [
|
||||
{
|
||||
"owner": "mypy",
|
||||
"pattern": [
|
||||
{
|
||||
"regexp": "^(.+):(\\d+):\\s(error|warning):\\s(.+)$",
|
||||
"file": 1,
|
||||
"line": 2,
|
||||
"severity": 3,
|
||||
"message": 4
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -1,173 +0,0 @@
|
||||
name: PR Test
|
||||
|
||||
on:
|
||||
push:
|
||||
branches: [main]
|
||||
paths:
|
||||
- "fastvideo/**/*.py"
|
||||
- ".github/workflows/pr-test.yml"
|
||||
pull_request:
|
||||
branches: [main]
|
||||
types: [opened, ready_for_review, synchronize, reopened]
|
||||
paths:
|
||||
- "fastvideo/**/*.py"
|
||||
- ".github/workflows/pr-test.yml"
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
run_encoder_test:
|
||||
description: "Run encoder-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_ssim_test:
|
||||
description: "Run ssim-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
|
||||
concurrency:
|
||||
group: pr-test-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
uses: ./.github/workflows/pre-commit.yml
|
||||
|
||||
change-filter:
|
||||
runs-on: ubuntu-latest
|
||||
needs: pre-commit
|
||||
if: ${{ github.event.pull_request.draft == false || github.event_name == 'workflow_dispatch' }}
|
||||
outputs:
|
||||
encoder-test: ${{ steps.filter.outputs.encoder-test }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
id: filter
|
||||
with:
|
||||
filters: |
|
||||
encoder-test:
|
||||
- 'fastvideo/v1/models/encoders/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/encoders/**'
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.encoder-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_encoder_test == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
environment: runpod-runners
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
- name: Set up SSH key
|
||||
run: |
|
||||
mkdir -p ~/.ssh
|
||||
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
|
||||
chmod 600 ~/.ssh/id_rsa
|
||||
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
|
||||
|
||||
- name: Install dependencies
|
||||
run: pip install requests
|
||||
|
||||
- name: Run tests on RunPod
|
||||
env:
|
||||
JOB_ID: "encoder-test"
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
timeout-minutes: 30
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
--gpu-type "NVIDIA A40"
|
||||
--gpu-count 1
|
||||
--volume-size 100
|
||||
--test-command "pip install -e .[test] &&
|
||||
pip install flash-attn==2.7.0.post2 --no-build-isolation &&
|
||||
pytest ./fastvideo/v1/tests/encoders -s"
|
||||
|
||||
- name: Terminate RunPod Instances
|
||||
if: ${{ always() }}
|
||||
env:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
JOB_ID: "encoder-test"
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
ssim-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
|
||||
runs-on: ubuntu-latest
|
||||
environment: runpod-runners
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
- name: Set up SSH key
|
||||
run: |
|
||||
mkdir -p ~/.ssh
|
||||
echo "${{ secrets.RUNPOD_PRIVATE_KEY }}" > ~/.ssh/id_rsa
|
||||
chmod 600 ~/.ssh/id_rsa
|
||||
ssh-keygen -y -f ~/.ssh/id_rsa > ~/.ssh/id_rsa.pub
|
||||
|
||||
- name: Install dependencies
|
||||
run: pip install requests
|
||||
|
||||
- name: Run tests on RunPod
|
||||
env:
|
||||
JOB_ID: "ssim-test"
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
timeout-minutes: 30
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
--gpu-type "NVIDIA A40"
|
||||
--gpu-count 2
|
||||
--disk-size 100
|
||||
--volume-size 100
|
||||
--test-command "pip install -e .[test] &&
|
||||
pip install flash-attn==2.7.0.post2 --no-build-isolation &&
|
||||
pytest ./fastvideo/v1/tests/ssim -vs"
|
||||
|
||||
- name: Terminate RunPod Instances
|
||||
if: ${{ always() }}
|
||||
env:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
JOB_ID: "ssim-test"
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
runpod-cleanup:
|
||||
needs: [encoder-test, ssim-test] # Add other jobs to this list as you create them
|
||||
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
- name: Install dependencies
|
||||
run: pip install requests
|
||||
|
||||
- name: Cleanup all RunPod instances
|
||||
env:
|
||||
JOB_IDS: '["encoder-test", "ssim-test"]' # JSON array of job IDs
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
@@ -1,18 +0,0 @@
|
||||
name: pre-commit
|
||||
|
||||
on:
|
||||
workflow_call:
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/actionlint.json"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/mypy.json"
|
||||
- uses: pre-commit/action@v3.0.1
|
||||
with:
|
||||
extra_args: --all-files --hook-stage manual
|
||||
@@ -0,0 +1,50 @@
|
||||
name: ruff
|
||||
|
||||
on:
|
||||
# Trigger the workflow on push or pull request,
|
||||
# but only for the main branch
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- pyproject.toml
|
||||
- requirements-lint.txt
|
||||
- .github/workflows/matchers/ruff.json
|
||||
- .github/workflows/ruff.yml
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
# This workflow is only relevant when one of the following files changes.
|
||||
# However, we have github configured to expect and require this workflow
|
||||
# to run and pass before github with auto-merge a pull request. Until github
|
||||
# allows more flexible auto-merge policy, we can just run this on every PR.
|
||||
# It doesn't take that long to run, anyway.
|
||||
#paths:
|
||||
# - "**/*.py"
|
||||
# - pyproject.toml
|
||||
# - requirements-lint.txt
|
||||
# - .github/workflows/matchers/ruff.json
|
||||
# - .github/workflows/ruff.yml
|
||||
|
||||
jobs:
|
||||
ruff:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v3
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: '3.12' # or any version you need
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install -r requirements-lint.txt
|
||||
- name: Analysing the code with ruff
|
||||
run: |
|
||||
ruff check .
|
||||
- name: Run isort
|
||||
run: |
|
||||
isort . --check-only
|
||||
@@ -15,7 +15,7 @@ jobs:
|
||||
new-version: ${{ steps.check-version.outputs.new-version }}
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
uses: actions/checkout@v3
|
||||
with:
|
||||
fetch-depth: 2
|
||||
|
||||
@@ -43,7 +43,7 @@ jobs:
|
||||
build_wheels:
|
||||
name: Build Wheel
|
||||
needs: check-version-change
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' }}
|
||||
if: needs.check-version-change.outputs.version-changed == 'true'
|
||||
runs-on: ${{ matrix.os }}
|
||||
|
||||
strategy:
|
||||
@@ -144,8 +144,8 @@ jobs:
|
||||
|
||||
publish_package:
|
||||
name: Publish package
|
||||
needs: [build_wheels, check-version-change]
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' }}
|
||||
needs: [build_wheels]
|
||||
if: needs.check-version-change.outputs.version-changed == 'true'
|
||||
runs-on: ubuntu-22.04
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
@@ -11,10 +11,10 @@ jobs:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v4
|
||||
uses: actions/checkout@v3
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: '3.12' # or any version you need
|
||||
|
||||
@@ -23,9 +23,11 @@ jobs:
|
||||
python -m pip install --upgrade pip setuptools wheel
|
||||
pip install torch
|
||||
pip install packaging ninja
|
||||
# remove st-attn dependency because no cuda environment
|
||||
sed -i '/st_attn/d' pyproject.toml
|
||||
pip install -e .
|
||||
pip install pytest
|
||||
|
||||
- name: Run Pytest
|
||||
run: |
|
||||
pytest --ignore csrc/sliding_tile_attention/test
|
||||
pytest --ignore csrc/sliding_tile_attention/test
|
||||
@@ -0,0 +1,38 @@
|
||||
name: yapf
|
||||
|
||||
on:
|
||||
# Trigger the workflow on push or pull request,
|
||||
# but only for the main branch
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- .github/workflows/yapf.yml
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "**/*.py"
|
||||
- .github/workflows/yapf.yml
|
||||
|
||||
jobs:
|
||||
yapf:
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out repository
|
||||
uses: actions/checkout@v3
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v4
|
||||
with:
|
||||
python-version: '3.12' # or any version you need
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
python -m pip install --upgrade pip
|
||||
pip install yapf==0.32.0
|
||||
pip install toml==0.10.2
|
||||
- name: Running yapf
|
||||
run: |
|
||||
yapf --diff --recursive .
|
||||
+1
-1
@@ -1,4 +1,5 @@
|
||||
__pycache__
|
||||
*.mp4
|
||||
.ipynb_checkpoints
|
||||
*.pth
|
||||
UCF-101/
|
||||
@@ -10,7 +11,6 @@ wandb/
|
||||
*.jpg
|
||||
*.safetensors
|
||||
*.mp4
|
||||
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
|
||||
*.png
|
||||
*.gif
|
||||
*.pth
|
||||
|
||||
@@ -1,80 +0,0 @@
|
||||
default_stages:
|
||||
- pre-commit # Run locally
|
||||
- manual # Run in CI
|
||||
exclude: |
|
||||
(?x)(
|
||||
fastvideo/v1/third_party/.*|
|
||||
csrc/.*|
|
||||
assets/.*|
|
||||
tests/.*|
|
||||
demo/.*|
|
||||
predict\.py|
|
||||
scripts/.*|
|
||||
fastvideo/data_preprocess/.*|
|
||||
fastvideo/dataset/.*|
|
||||
fastvideo/distill/.*|
|
||||
fastvideo/distill\.py|
|
||||
fastvideo/distill_adv\.py|
|
||||
fastvideo/models/.*|
|
||||
fastvideo/sample/.*|
|
||||
fastvideo/train\.py|
|
||||
fastvideo/utils/.*|
|
||||
.github/workflows/fastvideo-publish.yml|
|
||||
.github/workflows/sta-publish.yml
|
||||
)
|
||||
repos:
|
||||
- repo: https://github.com/google/yapf
|
||||
rev: v0.43.0
|
||||
hooks:
|
||||
- id: yapf
|
||||
args: [--in-place, --verbose]
|
||||
additional_dependencies: [toml] # TODO: Remove when yapf is upgraded
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
rev: v0.11.4
|
||||
hooks:
|
||||
- id: ruff
|
||||
args: [--output-format, github, --fix]
|
||||
- repo: https://github.com/codespell-project/codespell
|
||||
rev: v2.4.1
|
||||
hooks:
|
||||
- id: codespell
|
||||
additional_dependencies: ['tomli']
|
||||
args: ['--toml', 'pyproject.toml']
|
||||
- repo: https://github.com/PyCQA/isort
|
||||
rev: 0a0b7a830386ba6a31c2ec8316849ae4d1b8240d # 6.0.0
|
||||
hooks:
|
||||
- id: isort
|
||||
- repo: https://github.com/jackdewinter/pymarkdown
|
||||
rev: v0.9.29
|
||||
hooks:
|
||||
- id: pymarkdown
|
||||
args: [fix]
|
||||
- repo: https://github.com/rhysd/actionlint
|
||||
rev: v1.7.7
|
||||
hooks:
|
||||
- id: actionlint
|
||||
- repo: https://github.com/pre-commit/mirrors-mypy
|
||||
rev: v1.15.0
|
||||
hooks:
|
||||
- id: mypy
|
||||
args: [--python-version, '3.10', --follow-imports, "skip", ]
|
||||
additional_dependencies: [types-cachetools, types-setuptools, types-PyYAML, types-requests]
|
||||
- repo: local
|
||||
hooks:
|
||||
- id: check-filenames
|
||||
name: Check for spaces in all filenames
|
||||
entry: bash
|
||||
args:
|
||||
- -c
|
||||
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
|
||||
language: system
|
||||
always_run: true
|
||||
pass_filenames: false
|
||||
# Keep `suggestion` last
|
||||
- id: suggestion
|
||||
name: Suggestion
|
||||
entry: bash -c 'echo "To bypass pre-commit hooks, add --no-verify to git commit."'
|
||||
language: system
|
||||
verbose: true
|
||||
pass_filenames: false
|
||||
# Insert new entries above the `suggestion` entry
|
||||
@@ -1,21 +0,0 @@
|
||||
# Read the Docs configuration file
|
||||
# See https://docs.readthedocs.io/en/stable/config-file/v2.html for details
|
||||
|
||||
version: 2
|
||||
|
||||
build:
|
||||
os: ubuntu-22.04
|
||||
tools:
|
||||
python: "3.12"
|
||||
|
||||
sphinx:
|
||||
configuration: docs/source/conf.py
|
||||
fail_on_warning: true
|
||||
|
||||
# If using Sphinx, optionally build your docs in additional formats such as PDF
|
||||
formats: []
|
||||
|
||||
# Optionally declare the Python requirements required to build your docs
|
||||
python:
|
||||
install:
|
||||
- requirements: docs/requirements-docs.txt
|
||||
@@ -4,12 +4,19 @@
|
||||
|
||||
FastVideo is a lightweight framework for accelerating large video diffusion models.
|
||||
|
||||
|
||||
<p align="center">
|
||||
🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank">FastHunyuan</a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank">FastMochi</a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> Slack </a>
|
||||
</p>
|
||||
🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank">FastHunyuan</a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank">FastMochi</a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> Slack </a>
|
||||
</p>
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1
|
||||
|
||||
|
||||
|
||||
FastVideo currently offers: (with more to come)
|
||||
|
||||
- [NEW!] [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
@@ -21,6 +28,8 @@ FastVideo currently offers: (with more to come)
|
||||
|
||||
Dev in progress and highly experimental.
|
||||
|
||||
|
||||
|
||||
## 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/).
|
||||
@@ -28,32 +37,21 @@ Dev in progress and highly experimental.
|
||||
- ```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
|
||||
|
||||
## 🔧 Installation
|
||||
The code is tested on Python 3.10.0, CUDA 12.4 and H100.
|
||||
|
||||
```
|
||||
# Clone FastVideo
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
|
||||
# Install FastVideo
|
||||
pip install -e .
|
||||
|
||||
# Install Flash Attention (optional)
|
||||
pip install flash-attn==2.7.0.post2
|
||||
./env_setup.sh fastvideo
|
||||
```
|
||||
|
||||
To try Sliding Tile Attention (optional), please follow the instruction in [csrc/sliding_tile_attention/README.md](csrc/sliding_tile_attention/README.md) to install STA.
|
||||
|
||||
## 🚀 Inference
|
||||
### Inference StepVideo with Sliding Tile Attention
|
||||
### 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
|
||||
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
|
||||
@@ -61,49 +59,44 @@ 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
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model
|
||||
```
|
||||
|
||||
We provide two examples in the following script to run inference with STA + [TeaCache](https://github.com/ali-vilab/TeaCache) and STA only.
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_hunyuan_STA.sh
|
||||
```
|
||||
|
||||
### Video Demos using STA + Teacache
|
||||
Visit our [demo website](https://fast-video.github.io/) to explore our complete collection of examples. We shorten a single video generation process from 945s to 317s on H100.
|
||||
|
||||
### Inference FastHunyuan on single RTX4090
|
||||
We now support NF4 and LLM-INT8 quantized inference using BitsAndBytes for FastHunyuan. With NF4 quantization, inference can be performed on a single RTX 4090 GPU, requiring just 20GB of VRAM.
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan-diffusers --local_dir=data/FastHunyuan-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan_hf_quantization.sh
|
||||
```
|
||||
|
||||
For more information about the VRAM requirements for BitsAndBytes quantization, please refer to the table below (timing measured on an H100 GPU):
|
||||
|
||||
|
||||
| Configuration | Memory to Init Transformer | Peak Memory After Init Pipeline (Denoise) | Diffusion Time | End-to-End Time |
|
||||
|--------------------------------|----------------------------|--------------------------------------------|----------------|-----------------|
|
||||
| BF16 + Pipeline CPU Offload | 23.883G | 33.744G | 81s | 121.5s |
|
||||
| INT8 + Pipeline CPU Offload | 13.911G | 27.979G | 88s | 116.7s |
|
||||
| NF4 + Pipeline CPU Offload | 9.453G | 19.26G | 78s | 114.5s |
|
||||
|
||||
|
||||
|
||||
For improved quality in generated videos, we recommend using a GPU with 80GB of memory to run the BF16 model with the original Hunyuan pipeline. To execute the inference, use the following section:
|
||||
|
||||
### FastHunyuan
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan --local_dir=data/FastHunyuan --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan.sh
|
||||
```
|
||||
|
||||
You can also inference FastHunyuan in the [official Hunyuan github](https://github.com/Tencent/HunyuanVideo).
|
||||
|
||||
### FastMochi
|
||||
@@ -115,99 +108,79 @@ python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastMochi-diffuser
|
||||
bash scripts/inference/inference_mochi_sp.sh
|
||||
```
|
||||
|
||||
|
||||
## 🎯 Distill
|
||||
Our distillation recipe is based on [Phased Consistency Model](https://github.com/G-U-N/Phased-Consistency-Model). We did not find significant improvement using multi-phase distillation, so we keep the one phase setup similar to the original latent consistency model's recipe.
|
||||
We use the [MixKit](https://huggingface.co/datasets/LanguageBind/Open-Sora-Plan-v1.1.0/tree/main/all_mixkit) dataset for distillation. To avoid running the text encoder and VAE during training, we preprocess all data to generate text embeddings and VAE latents.
|
||||
Preprocessing instructions can be found [data_preprocess.md](docs/data_preprocess.md). For convenience, we also provide preprocessed data that can be downloaded directly using the following command:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/HD-Mixkit-Finetune-Hunyuan --local_dir=data/HD-Mixkit-Finetune-Hunyuan --repo_type=dataset
|
||||
```
|
||||
|
||||
Next, download the original model weights with:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model # original hunyuan
|
||||
python scripts/huggingface/download_hf.py --repo_id=genmo/mochi-1-preview --local_dir=data/mochi --repo_type=model # original mochi
|
||||
```
|
||||
|
||||
To launch the distillation process, use the following commands:
|
||||
|
||||
```
|
||||
bash scripts/distill/distill_hunyuan.sh # for hunyuan
|
||||
bash scripts/distill/distill_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
We also provide an optional script for distillation with adversarial loss, located at `fastvideo/distill_adv.py`. Although we tried adversarial loss, we did not observe significant improvements.
|
||||
## Finetune
|
||||
### ⚡ Full Finetune
|
||||
Ensure your data is prepared and preprocessed in the format specified in [data_preprocess.md](docs/data_preprocess.md). For convenience, we also provide a mochi preprocessed Black Myth Wukong data that can be downloaded directly:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Mochi-Black-Myth --local_dir=data/Mochi-Black-Myth --repo_type=dataset
|
||||
```
|
||||
|
||||
Download the original model weights as specified in [Distill Section](#-distill):
|
||||
|
||||
Then you can run the finetune with:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
### ⚡ Lora Finetune
|
||||
### ⚡ Lora Finetune
|
||||
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight --local_dir=data/Hunyuan-Black-Myth-Wukong-lora-weight --repo_type=model
|
||||
```
|
||||
|
||||
#### Minimum Hardware Requirement
|
||||
- 40 GB GPU memory each for 2 GPUs with lora.
|
||||
- 30 GB GPU memory each for 2 GPUs with CPU offload and lora.
|
||||
- 30 GB GPU memory each for 2 GPUs with CPU offload and lora.
|
||||
|
||||
|
||||
Currently, both Mochi and Hunyuan models support Lora finetuning through diffusers. To generate personalized videos from your own dataset, you'll need to follow three main steps: dataset preparation, finetuning, and inference.
|
||||
|
||||
#### Dataset Preparation
|
||||
We provide scripts to better help you get started to train on your own characters!
|
||||
We provide scripts to better help you get started to train on your own characters!
|
||||
You can run this to organize your dataset to get the videos2caption.json before preprocess. Specify your video folder and corresponding caption folder (caption files should be .txt files and have the same name with its video):
|
||||
|
||||
```
|
||||
python scripts/dataset_preparation/prepare_json_file.py --video_dir data/input_videos/ --prompt_dir data/captions/ --output_path data/output_folder/videos2caption.json --verbose
|
||||
```
|
||||
|
||||
Also, we provide script to resize your videos:
|
||||
|
||||
```
|
||||
python scripts/data_preprocess/resize_videos.py
|
||||
python scripts/data_preprocess/resize_videos.py
|
||||
```
|
||||
|
||||
#### Finetuning
|
||||
After basic dataset preparation and preprocess, you can start to finetune your model using Lora:
|
||||
|
||||
```
|
||||
bash scripts/finetune/finetune_hunyuan_hf_lora.sh
|
||||
```
|
||||
|
||||
#### Inference
|
||||
For inference with Lora checkpoint, you can run the following scripts with additional parameter `--lora_checkpoint_dir`:
|
||||
|
||||
```
|
||||
bash scripts/inference/inference_hunyuan_hf.sh
|
||||
bash scripts/inference/inference_hunyuan_hf.sh
|
||||
```
|
||||
|
||||
**We also provide scripts for Mochi in the same directory.**
|
||||
|
||||
#### Finetune with Both Image and Video
|
||||
Our codebase support finetuning with both image and video.
|
||||
|
||||
Our codebase support finetuning with both image and video.
|
||||
```bash
|
||||
bash scripts/finetune/finetune_hunyuan.sh
|
||||
bash scripts/finetune/finetune_mochi_lora_mix.sh
|
||||
```
|
||||
|
||||
For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` option in your script.
|
||||
|
||||
## 📑 Development Plan
|
||||
@@ -232,26 +205,26 @@ We learned and reused code from the following projects: [PCM](https://github.com
|
||||
|
||||
We thank MBZUAI and Anyscale for their support throughout this project.
|
||||
|
||||
## Citation
|
||||
## Citation
|
||||
If you use FastVideo for your research, please cite our paper:
|
||||
|
||||
```bibtex
|
||||
@misc{zhang2025fastvideogenerationsliding,
|
||||
title={Fast Video Generation with Sliding Tile Attention},
|
||||
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},
|
||||
url={https://arxiv.org/abs/2502.04507},
|
||||
}
|
||||
@misc{ding2025efficientvditefficientvideodiffusion,
|
||||
title={Efficient-vDiT: Efficient Video Diffusion Transformers With Attention Tile},
|
||||
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},
|
||||
url={https://arxiv.org/abs/2502.06155},
|
||||
}
|
||||
```
|
||||
|
||||
@@ -8,7 +8,7 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
|
||||
seq_length = q_all.shape[2]
|
||||
if has_text:
|
||||
assert q_all.shape[
|
||||
2] >= 115200, "STA currently only supports video with latent size (30, 48, 80), which is 117 frames x 768 x 1280 pixels"
|
||||
2] == 115456, "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
|
||||
|
||||
@@ -1,24 +0,0 @@
|
||||
# Minimal makefile for Sphinx documentation
|
||||
#
|
||||
|
||||
# You can set these variables from the command line, and also
|
||||
# from the environment for the first two.
|
||||
SPHINXOPTS ?=
|
||||
SPHINXBUILD ?= sphinx-build
|
||||
SOURCEDIR = source
|
||||
BUILDDIR = build
|
||||
|
||||
# Put it first so that "make" without argument is like "make help".
|
||||
help:
|
||||
@$(SPHINXBUILD) -M help "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
|
||||
|
||||
.PHONY: help Makefile
|
||||
|
||||
# Catch-all target: route all unknown targets to Sphinx using the new
|
||||
# "make mode" option. $(O) is meant as a shortcut for $(SPHINXOPTS).
|
||||
%: Makefile
|
||||
@$(SPHINXBUILD) -M $@ "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
|
||||
|
||||
clean:
|
||||
@$(SPHINXBUILD) -M clean "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
|
||||
rm -rf "$(SOURCEDIR)/getting_started/examples"
|
||||
@@ -1,20 +0,0 @@
|
||||
# FastVideo documents
|
||||
|
||||
## Build the docs
|
||||
|
||||
```bash
|
||||
# Install dependencies.
|
||||
pip install -r requirements-docs.txt
|
||||
|
||||
# Build the docs.
|
||||
make clean
|
||||
make html
|
||||
```
|
||||
|
||||
## Open the docs with your browser
|
||||
|
||||
```bash
|
||||
python -m http.server -d build/html/
|
||||
```
|
||||
|
||||
Launch your browser and open localhost:8000.
|
||||
@@ -1,16 +1,16 @@
|
||||
|
||||
|
||||
|
||||
## 🧱 Data Preprocess
|
||||
|
||||
To save GPU memory, we precompute text embeddings and VAE latents to eliminate the need to load the text encoder and VAE during training.
|
||||
|
||||
We provide a sample dataset to help you get started. Download the source media using the following command:
|
||||
|
||||
We provide a sample dataset to help you get started. Download the source media using the following command:
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Image-Vid-Finetune-Src --local_dir=data/Image-Vid-Finetune-Src --repo_type=dataset
|
||||
```
|
||||
|
||||
To preprocess the dataset for fine-tuning or distillation, run:
|
||||
|
||||
```
|
||||
bash scripts/preprocess/preprocess_mochi_data.sh # for mochi
|
||||
bash scripts/preprocess/preprocess_hunyuan_data.sh # for hunyuan
|
||||
@@ -33,16 +33,13 @@ path_to_dataset_folder/
|
||||
Format the JSON file as a list, where each item represents a media source:
|
||||
|
||||
For image media,
|
||||
|
||||
```
|
||||
{
|
||||
"path": "0.jpg",
|
||||
"cap": ["captions"]
|
||||
}
|
||||
```
|
||||
|
||||
For video media,
|
||||
|
||||
For video media,
|
||||
```
|
||||
{
|
||||
"path": "1.mp4",
|
||||
@@ -65,9 +62,7 @@ path_to_media_source_foder,path_to_json_file
|
||||
```
|
||||
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/preprocess_****_data.sh` accordingly and run:
|
||||
|
||||
```
|
||||
bash scripts/preprocess/preprocess_****_data.sh
|
||||
```
|
||||
|
||||
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
|
||||
|
||||
@@ -1,35 +0,0 @@
|
||||
@ECHO OFF
|
||||
|
||||
pushd %~dp0
|
||||
|
||||
REM Command file for Sphinx documentation
|
||||
|
||||
if "%SPHINXBUILD%" == "" (
|
||||
set SPHINXBUILD=sphinx-build
|
||||
)
|
||||
set SOURCEDIR=source
|
||||
set BUILDDIR=build
|
||||
|
||||
%SPHINXBUILD% >NUL 2>NUL
|
||||
if errorlevel 9009 (
|
||||
echo.
|
||||
echo.The 'sphinx-build' command was not found. Make sure you have Sphinx
|
||||
echo.installed, then set the SPHINXBUILD environment variable to point
|
||||
echo.to the full path of the 'sphinx-build' executable. Alternatively you
|
||||
echo.may add the Sphinx directory to PATH.
|
||||
echo.
|
||||
echo.If you don't have Sphinx installed, grab it from
|
||||
echo.https://www.sphinx-doc.org/
|
||||
exit /b 1
|
||||
)
|
||||
|
||||
if "%1" == "" goto help
|
||||
|
||||
%SPHINXBUILD% -M %1 %SOURCEDIR% %BUILDDIR% %SPHINXOPTS% %O%
|
||||
goto end
|
||||
|
||||
:help
|
||||
%SPHINXBUILD% -M help %SOURCEDIR% %BUILDDIR% %SPHINXOPTS% %O%
|
||||
|
||||
:end
|
||||
popd
|
||||
@@ -1,25 +0,0 @@
|
||||
sphinx==6.2.1
|
||||
sphinx-argparse==0.4.0
|
||||
sphinx-book-theme==1.0.1
|
||||
sphinx-copybutton==0.5.2
|
||||
sphinx-design==0.6.1
|
||||
sphinx-togglebutton==0.3.2
|
||||
myst-parser==3.0.1
|
||||
msgspec
|
||||
cloudpickle
|
||||
|
||||
# packages to install to build the documentation
|
||||
cachetools
|
||||
pydantic >= 2.8
|
||||
-f https://download.pytorch.org/whl/cpu
|
||||
torch
|
||||
py-cpuinfo
|
||||
transformers
|
||||
mistral_common >= 1.5.4
|
||||
aiohttp
|
||||
starlette
|
||||
openai # Required by docs/source/serving/openai_compatible_server.md's vllm.entrypoints.openai.cli_args
|
||||
fastapi # Required by docs/source/serving/openai_compatible_server.md's vllm.entrypoints.openai.cli_args
|
||||
partial-json-parser # Required by docs/source/serving/openai_compatible_server.md's vllm.entrypoints.openai.cli_args
|
||||
requests
|
||||
zmq
|
||||
@@ -1,51 +0,0 @@
|
||||
# Seed Parameter Behavior in vLLM
|
||||
|
||||
## Overview
|
||||
|
||||
The `seed` parameter in vLLM is used to control the random states for various random number generators. This parameter can affect the behavior of random operations in user code, especially when working with models in vLLM.
|
||||
|
||||
## Default Behavior
|
||||
|
||||
By default, the `seed` parameter is set to `None`. When the `seed` parameter is `None`, the global random states for `random`, `np.random`, and `torch.manual_seed` are not set. This means that the random operations will behave as expected, without any fixed random states.
|
||||
|
||||
## Specifying a Seed
|
||||
|
||||
If a specific seed value is provided, the global random states for `random`, `np.random`, and `torch.manual_seed` will be set accordingly. This can be useful for reproducibility, as it ensures that the random operations produce the same results across multiple runs.
|
||||
|
||||
## Example Usage
|
||||
|
||||
### Without Specifying a Seed
|
||||
|
||||
```python
|
||||
import random
|
||||
from vllm import LLM
|
||||
|
||||
# Initialize a vLLM model without specifying a seed
|
||||
model = LLM(model="Qwen/Qwen2.5-0.5B-Instruct")
|
||||
|
||||
# Try generating random numbers
|
||||
print(random.randint(0, 100)) # Outputs different numbers across runs
|
||||
```
|
||||
|
||||
### Specifying a Seed
|
||||
|
||||
```python
|
||||
import random
|
||||
from vllm import LLM
|
||||
|
||||
# Initialize a vLLM model with a specific seed
|
||||
model = LLM(model="Qwen/Qwen2.5-0.5B-Instruct", seed=42)
|
||||
|
||||
# Try generating random numbers
|
||||
print(random.randint(0, 100)) # Outputs the same number across runs
|
||||
```
|
||||
|
||||
## Important Notes
|
||||
|
||||
- If the `seed` parameter is not specified, the behavior of global random states remains unaffected.
|
||||
- If a specific seed value is provided, the global random states for `random`, `np.random`, and `torch.manual_seed` will be set to that value.
|
||||
- This behavior can be useful for reproducibility but may lead to non-intuitive behavior if the user is not explicitly aware of it.
|
||||
|
||||
## Conclusion
|
||||
|
||||
Understanding the behavior of the `seed` parameter in vLLM is crucial for ensuring the expected behavior of random operations in your code. By default, the `seed` parameter is set to `None`, which means that the global random states are not affected. However, specifying a seed value can help achieve reproducibility in your experiments.
|
||||
@@ -1,8 +0,0 @@
|
||||
.vertical-table-header th.head:not(.stub) {
|
||||
writing-mode: sideways-lr;
|
||||
white-space: nowrap;
|
||||
max-width: 0;
|
||||
p {
|
||||
margin: 0;
|
||||
}
|
||||
}
|
||||
@@ -1,18 +0,0 @@
|
||||
// Update URL search params when tab is clicked
|
||||
document.addEventListener("DOMContentLoaded", function () {
|
||||
const tabs = document.querySelectorAll(".sd-tab-label");
|
||||
|
||||
function updateURL(tab) {
|
||||
const syncGroup = tab.getAttribute("data-sync-group");
|
||||
const syncId = tab.getAttribute("data-sync-id");
|
||||
if (syncGroup && syncId) {
|
||||
const url = new URL(window.location);
|
||||
url.searchParams.set(syncGroup, syncId);
|
||||
window.history.replaceState(null, "", url);
|
||||
}
|
||||
}
|
||||
|
||||
tabs.forEach(tab => {
|
||||
tab.addEventListener("click", () => updateURL(tab));
|
||||
});
|
||||
});
|
||||
@@ -1,39 +0,0 @@
|
||||
<style>
|
||||
.notification-bar {
|
||||
width: 100vw;
|
||||
display: flex;
|
||||
justify-content: center;
|
||||
align-items: center;
|
||||
font-size: 16px;
|
||||
padding: 0 6px 0 6px;
|
||||
}
|
||||
.notification-bar p {
|
||||
margin: 0;
|
||||
}
|
||||
.notification-bar a {
|
||||
font-weight: bold;
|
||||
text-decoration: none;
|
||||
}
|
||||
|
||||
/* Light mode styles (default) */
|
||||
.notification-bar {
|
||||
background-color: #fff3cd;
|
||||
color: #856404;
|
||||
}
|
||||
.notification-bar a {
|
||||
color: #d97706;
|
||||
}
|
||||
|
||||
/* Dark mode styles */
|
||||
html[data-theme=dark] .notification-bar {
|
||||
background-color: #333;
|
||||
color: #ddd;
|
||||
}
|
||||
html[data-theme=dark] .notification-bar a {
|
||||
color: #ffa500; /* Brighter color for visibility */
|
||||
}
|
||||
</style>
|
||||
|
||||
<div class="notification-bar">
|
||||
<p>You are viewing the latest developer preview docs. <a href="https://docs.vllm.ai/en/stable/">Click here</a> to view docs for the latest stable release.</p>
|
||||
</div>
|
||||
@@ -1,260 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Configuration file for the Sphinx documentation builder.
|
||||
#
|
||||
# This file only contains a selection of the most common options. For a full
|
||||
# list see the documentation:
|
||||
# https://www.sphinx-doc.org/en/master/usage/configuration.html
|
||||
|
||||
# -- Path setup --------------------------------------------------------------
|
||||
|
||||
# If extensions (or modules to document with autodoc) are in another directory,
|
||||
# add these directories to sys.path here. If the directory is relative to the
|
||||
# documentation root, use os.path.abspath to make it absolute, like shown here.
|
||||
|
||||
import datetime
|
||||
import inspect
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
from typing import Optional
|
||||
|
||||
import requests
|
||||
from sphinx.ext import autodoc
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
sys.path.append(os.path.abspath("../.."))
|
||||
|
||||
# -- Project information -----------------------------------------------------
|
||||
|
||||
project = 'FastVideo'
|
||||
copyright = f'{datetime.datetime.now().year}, FastVideo Team'
|
||||
author = 'the FastVideo Team'
|
||||
|
||||
# -- General configuration ---------------------------------------------------
|
||||
|
||||
# Add any Sphinx extension module names here, as strings. They can be
|
||||
# extensions coming with Sphinx (named 'sphinx.ext.*') or your custom
|
||||
# ones.
|
||||
extensions = [
|
||||
"sphinx.ext.napoleon",
|
||||
"sphinx.ext.linkcode",
|
||||
"sphinx.ext.intersphinx",
|
||||
"sphinx_copybutton",
|
||||
"sphinx.ext.autodoc",
|
||||
"sphinx.ext.autosummary",
|
||||
"myst_parser",
|
||||
"sphinxarg.ext",
|
||||
"sphinx_design",
|
||||
"sphinx_togglebutton",
|
||||
]
|
||||
myst_enable_extensions = [
|
||||
"colon_fence",
|
||||
]
|
||||
|
||||
# Add any paths that contain templates here, relative to this directory.
|
||||
templates_path = ['_templates']
|
||||
|
||||
# List of patterns, relative to source directory, that match files and
|
||||
# directories to ignore when looking for source files.
|
||||
# This pattern also affects html_static_path and html_extra_path.
|
||||
exclude_patterns: list[str] = ["**/*.template.md", "**/*.inc.md"]
|
||||
|
||||
# Exclude the prompt "$" when copying code
|
||||
copybutton_prompt_text = r"\$ "
|
||||
copybutton_prompt_is_regexp = True
|
||||
|
||||
# -- Options for HTML output -------------------------------------------------
|
||||
|
||||
# The theme to use for HTML and HTML Help pages. See the documentation for
|
||||
# a list of builtin themes.
|
||||
#
|
||||
html_title = project
|
||||
html_theme = 'sphinx_book_theme'
|
||||
html_logo = '../../assets/logo.jpg'
|
||||
#html_favicon = 'assets/logos/vllm-logo-only-light.ico'
|
||||
html_theme_options = {
|
||||
'path_to_docs': 'docs/source',
|
||||
'repository_url': 'https://github.com/hao-ai-lab/FastVideo/',
|
||||
'use_repository_button': True,
|
||||
'use_edit_page_button': True,
|
||||
}
|
||||
# Add any paths that contain custom static files (such as style sheets) here,
|
||||
# relative to this directory. They are copied after the builtin static files,
|
||||
# so a file named "default.css" will overwrite the builtin "default.css".
|
||||
html_static_path = ["_static"]
|
||||
html_js_files = ["custom.js"]
|
||||
html_css_files = ["custom.css"]
|
||||
|
||||
myst_url_schemes = {
|
||||
'http': None,
|
||||
'https': None,
|
||||
'mailto': None,
|
||||
'ftp': None,
|
||||
"gh-issue": {
|
||||
"url":
|
||||
"https://github.com/hao-ai-lab/FastVideo/issues/{{path}}#{{fragment}}",
|
||||
"title": "Issue #{{path}}",
|
||||
"classes": ["github"],
|
||||
},
|
||||
"gh-pr": {
|
||||
"url":
|
||||
"https://github.com/hao-ai-lab/FastVideo/pull/{{path}}#{{fragment}}",
|
||||
"title": "Pull Request #{{path}}",
|
||||
"classes": ["github"],
|
||||
},
|
||||
"gh-dir": {
|
||||
"url": "https://github.com/hao-ai-lab/FastVideo/tree/main/{{path}}",
|
||||
"title": "{{path}}",
|
||||
"classes": ["github"],
|
||||
},
|
||||
"gh-file": {
|
||||
"url": "https://github.com/hao-ai-lab/FastVideo/blob/main/{{path}}",
|
||||
"title": "{{path}}",
|
||||
"classes": ["github"],
|
||||
},
|
||||
}
|
||||
|
||||
# see https://docs.readthedocs.io/en/stable/reference/environment-variables.html # noqa
|
||||
READTHEDOCS_VERSION_TYPE = os.environ.get('READTHEDOCS_VERSION_TYPE')
|
||||
if READTHEDOCS_VERSION_TYPE == "tag":
|
||||
# remove the warning banner if the version is a tagged release
|
||||
header_file = os.path.join(os.path.dirname(__file__),
|
||||
"_templates/sections/header.html")
|
||||
# The file might be removed already if the build is triggered multiple times
|
||||
# (readthedocs build both HTML and PDF versions separately)
|
||||
if os.path.exists(header_file):
|
||||
os.remove(header_file)
|
||||
|
||||
|
||||
# Generate additional rst documentation here.
|
||||
def setup(app):
|
||||
from docs.source.generate_examples import generate_examples
|
||||
generate_examples()
|
||||
|
||||
|
||||
_cached_base: str = ""
|
||||
_cached_branch: str = ""
|
||||
|
||||
|
||||
def get_repo_base_and_branch(
|
||||
pr_number: str) -> tuple[Optional[str], Optional[str]]:
|
||||
global _cached_base, _cached_branch
|
||||
if _cached_base and _cached_branch:
|
||||
return _cached_base, _cached_branch
|
||||
|
||||
url = f"https://api.github.com/repos/hao-ai-lab/FastVideo/pulls/{pr_number}"
|
||||
response = requests.get(url)
|
||||
if response.status_code == 200:
|
||||
data = response.json()
|
||||
_cached_base = data['head']['repo']['full_name']
|
||||
_cached_branch = data['head']['ref']
|
||||
return _cached_base, _cached_branch
|
||||
else:
|
||||
logger.error("Failed to fetch PR details: %s", response)
|
||||
return None, None
|
||||
|
||||
|
||||
def linkcode_resolve(domain, info):
|
||||
if domain != 'py':
|
||||
return None
|
||||
if not info['module']:
|
||||
return None
|
||||
module = info['module']
|
||||
|
||||
# try to determine the correct file and line number to link to
|
||||
obj = sys.modules[module]
|
||||
|
||||
# get as specific as we can
|
||||
lineno: int = 0
|
||||
filename: str = ""
|
||||
try:
|
||||
for part in info['fullname'].split('.'):
|
||||
obj = getattr(obj, part)
|
||||
|
||||
if not (inspect.isclass(obj) or inspect.isfunction(obj)
|
||||
or inspect.ismethod(obj)):
|
||||
obj = obj.__class__ # type: ignore[assignment]
|
||||
|
||||
lineno = inspect.getsourcelines(obj)[1]
|
||||
filename = (inspect.getsourcefile(obj)
|
||||
or f"{filename}.py").split("FastVideo/", 1)[1]
|
||||
except Exception:
|
||||
# For some things, like a class member, won't work, so
|
||||
# we'll use the line number of the parent (the class)
|
||||
pass
|
||||
|
||||
if filename.startswith("checkouts/"):
|
||||
# a PR build on readthedocs
|
||||
pr_number = filename.split("/")[1]
|
||||
filename = filename.split("/", 2)[2]
|
||||
base, branch = get_repo_base_and_branch(pr_number)
|
||||
if base and branch:
|
||||
return f"https://github.com/{base}/blob/{branch}/{filename}#L{lineno}"
|
||||
|
||||
# Otherwise, link to the source file on the main branch
|
||||
return f"https://github.com/hao-ai-lab/FastVideo/blob/main/{filename}#L{lineno}"
|
||||
|
||||
|
||||
# Mock out external dependencies here, otherwise the autodoc pages may be blank.
|
||||
autodoc_mock_imports = [
|
||||
"blake3",
|
||||
"compressed_tensors",
|
||||
"cpuinfo",
|
||||
"cv2",
|
||||
"torch",
|
||||
"transformers",
|
||||
"psutil",
|
||||
"prometheus_client",
|
||||
"sentencepiece",
|
||||
"vllm._C",
|
||||
"PIL",
|
||||
"numpy",
|
||||
'triton',
|
||||
"tqdm",
|
||||
"tensorizer",
|
||||
"pynvml",
|
||||
"outlines",
|
||||
"xgrammar",
|
||||
"librosa",
|
||||
"soundfile",
|
||||
"gguf",
|
||||
"lark",
|
||||
"decord",
|
||||
]
|
||||
|
||||
for mock_target in autodoc_mock_imports:
|
||||
if mock_target in sys.modules:
|
||||
logger.info(
|
||||
"Potentially problematic mock target (%s) found; "
|
||||
"autodoc_mock_imports cannot mock modules that have already "
|
||||
"been loaded into sys.modules when the sphinx build starts.",
|
||||
mock_target)
|
||||
|
||||
|
||||
class MockedClassDocumenter(autodoc.ClassDocumenter):
|
||||
"""Remove note about base class when a class is derived from object."""
|
||||
|
||||
def add_line(self, line: str, source: str, *lineno: int) -> None:
|
||||
if line == " Bases: :py:class:`object`":
|
||||
return
|
||||
super().add_line(line, source, *lineno)
|
||||
|
||||
|
||||
autodoc.ClassDocumenter = MockedClassDocumenter
|
||||
|
||||
intersphinx_mapping = {
|
||||
"python": ("https://docs.python.org/3", None),
|
||||
"typing_extensions":
|
||||
("https://typing-extensions.readthedocs.io/en/latest", None),
|
||||
"aiohttp": ("https://docs.aiohttp.org/en/stable", None),
|
||||
"pillow": ("https://pillow.readthedocs.io/en/stable", None),
|
||||
"numpy": ("https://numpy.org/doc/stable", None),
|
||||
"torch": ("https://pytorch.org/docs/stable", None),
|
||||
"psutil": ("https://psutil.readthedocs.io/en/stable", None),
|
||||
}
|
||||
|
||||
autodoc_preserve_defaults = True
|
||||
autodoc_warningiserror = True
|
||||
|
||||
navigation_with_keys = False
|
||||
@@ -1,50 +0,0 @@
|
||||
# Contributing to FastVideo
|
||||
|
||||
Thank you for your interest in contributing to FastVideo. We want to make the process as smooth for you as possible and this is a guide to help get you started!
|
||||
|
||||
Our community is open to everyone and welcomes any contributions no matter how large or small.
|
||||
|
||||
# Developer Environment:
|
||||
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only support Linux and CUDA GPUs, but we hope to support other platforms in the future.
|
||||
|
||||
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
|
||||
|
||||
Install Miniconda:
|
||||
|
||||
```
|
||||
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
|
||||
bash Miniconda3-latest-Linux-x86_64.sh
|
||||
source ~/.bashrc
|
||||
```
|
||||
|
||||
Create and activate a Conda environment for FastVideo:
|
||||
|
||||
```
|
||||
conda create -n fastvideo python=3.10 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
Clone the FastVideo repository and go to the FastVideo directory:
|
||||
|
||||
```
|
||||
git clone https://github.com/vllm-project/vllm.git && cd vllm
|
||||
|
||||
```
|
||||
|
||||
Now you can install FastVideo and setup git hooks for running linting. By using `pre-commit`, the linters will run and have to pass before you'll be able to make a commit.
|
||||
|
||||
```bash
|
||||
pip install -e .[dev]
|
||||
|
||||
# Can also install flash-attn (optional)
|
||||
pip install flash-attn==2.7.0.post2 --no-build-isolation
|
||||
|
||||
# Linting, formatting and static type checking
|
||||
pre-commit install --hook-type pre-commit --hook-type commit-msg
|
||||
|
||||
# You can manually run pre-commit with
|
||||
pre-commit run --all-files
|
||||
|
||||
# Unit tests
|
||||
pytest tests/
|
||||
```
|
||||
@@ -1,246 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import itertools
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
ROOT_DIR = Path(__file__).parent.parent.parent.resolve()
|
||||
ROOT_DIR_RELATIVE = '../../../..'
|
||||
EXAMPLE_DIR = ROOT_DIR / "examples"
|
||||
EXAMPLE_DOC_DIR = ROOT_DIR / "docs/source/getting_started/examples"
|
||||
|
||||
|
||||
def fix_case(text: str) -> str:
|
||||
subs = {
|
||||
"api": "API",
|
||||
"cli": "CLI",
|
||||
"cpu": "CPU",
|
||||
"llm": "LLM",
|
||||
"tpu": "TPU",
|
||||
"aqlm": "AQLM",
|
||||
"gguf": "GGUF",
|
||||
"lora": "LoRA",
|
||||
"rlhf": "RLHF",
|
||||
"vllm": "vLLM",
|
||||
"openai": "OpenAI",
|
||||
"multilora": "MultiLoRA",
|
||||
"mlpspeculator": "MLPSpeculator",
|
||||
r"fp\d+": lambda x: x.group(0).upper(), # e.g. fp16, fp32
|
||||
r"int\d+": lambda x: x.group(0).upper(), # e.g. int8, int16
|
||||
}
|
||||
for pattern, repl in subs.items():
|
||||
text = re.sub(rf'\b{pattern}\b', repl, text,
|
||||
flags=re.IGNORECASE) # type: ignore[call-overload]
|
||||
return text
|
||||
|
||||
|
||||
@dataclass
|
||||
class Index:
|
||||
"""
|
||||
Index class to generate a structured document index.
|
||||
|
||||
Attributes:
|
||||
path (Path): The path save the index file to.
|
||||
title (str): The title of the index.
|
||||
description (str): A brief description of the index.
|
||||
caption (str): An optional caption for the table of contents.
|
||||
maxdepth (int): The maximum depth of the table of contents. Defaults to 1.
|
||||
documents (list[str]): A list of document paths to include in the index. Defaults to an empty list.
|
||||
|
||||
Methods:
|
||||
generate() -> str:
|
||||
Generates the index content as a string in the specified format.
|
||||
""" # noqa: E501
|
||||
path: Path
|
||||
title: str
|
||||
description: str
|
||||
caption: str
|
||||
maxdepth: int = 1
|
||||
documents: list[str] = field(default_factory=list)
|
||||
|
||||
def generate(self) -> str:
|
||||
content = f"# {self.title}\n\n{self.description}\n\n"
|
||||
content += ":::{toctree}\n"
|
||||
content += f":caption: {self.caption}\n:maxdepth: {self.maxdepth}\n"
|
||||
content += "\n".join(self.documents) + "\n:::\n"
|
||||
return content
|
||||
|
||||
|
||||
@dataclass
|
||||
class Example:
|
||||
"""
|
||||
Example class for generating documentation content from a given path.
|
||||
|
||||
Attributes:
|
||||
path (Path): The path to the main directory or file.
|
||||
category (str): The category of the document.
|
||||
main_file (Path): The main file in the directory.
|
||||
other_files (list[Path]): list of other files in the directory.
|
||||
title (str): The title of the document.
|
||||
|
||||
Methods:
|
||||
__post_init__(): Initializes the main_file, other_files, and title attributes.
|
||||
determine_main_file() -> Path: Determines the main file in the given path.
|
||||
determine_other_files() -> list[Path]: Determines other files in the directory excluding the main file.
|
||||
determine_title() -> str: Determines the title of the document.
|
||||
generate() -> str: Generates the documentation content.
|
||||
""" # noqa: E501
|
||||
path: Path
|
||||
category: Optional[str] = None
|
||||
main_file: Path = field(init=False)
|
||||
other_files: list[Path] = field(init=False)
|
||||
title: str = field(init=False)
|
||||
|
||||
def __post_init__(self):
|
||||
self.main_file = self.determine_main_file()
|
||||
self.other_files = self.determine_other_files()
|
||||
self.title = self.determine_title()
|
||||
|
||||
def determine_main_file(self) -> Path:
|
||||
"""
|
||||
Determines the main file in the given path.
|
||||
If the path is a file, it returns the path itself. Otherwise, it searches
|
||||
for Markdown files (*.md) in the directory and returns the first one found.
|
||||
Returns:
|
||||
Path: The main file path, either the original path if it's a file or the first
|
||||
Markdown file found in the directory.
|
||||
Raises:
|
||||
IndexError: If no Markdown files are found in the directory.
|
||||
""" # noqa: E501
|
||||
return self.path if self.path.is_file() else list(
|
||||
self.path.glob("*.md")).pop()
|
||||
|
||||
def determine_other_files(self) -> list[Path]:
|
||||
"""
|
||||
Determine other files in the directory excluding the main file.
|
||||
|
||||
This method checks if the given path is a file. If it is, it returns an empty list.
|
||||
Otherwise, it recursively searches through the directory and returns a list of all
|
||||
files that are not the main file.
|
||||
|
||||
Returns:
|
||||
list[Path]: A list of Path objects representing the other files in the directory.
|
||||
""" # noqa: E501
|
||||
if self.path.is_file():
|
||||
return []
|
||||
is_other_file = lambda file: file.is_file() and file != self.main_file
|
||||
return [file for file in self.path.rglob("*")
|
||||
if is_other_file(file)] # type: ignore[no-untyped-call]
|
||||
|
||||
def determine_title(self) -> str:
|
||||
return fix_case(self.path.stem.replace("_", " ").title())
|
||||
|
||||
def generate(self) -> str:
|
||||
# Convert the path to a relative path from __file__
|
||||
make_relative = lambda path: ROOT_DIR_RELATIVE / path.relative_to(
|
||||
ROOT_DIR)
|
||||
|
||||
content = f"Source <gh-file:{self.path.relative_to(ROOT_DIR)}>.\n\n"
|
||||
include = "include" if self.main_file.suffix == ".md" else \
|
||||
"literalinclude"
|
||||
if include == "literalinclude":
|
||||
content += f"# {self.title}\n\n"
|
||||
content += f":::{{{include}}} {make_relative(self.main_file)}\n" # type: ignore[no-untyped-call]
|
||||
if include == "literalinclude":
|
||||
content += f":language: {self.main_file.suffix[1:]}\n"
|
||||
content += ":::\n\n"
|
||||
|
||||
if not self.other_files:
|
||||
return content
|
||||
|
||||
content += "## Example materials\n\n"
|
||||
for file in sorted(self.other_files):
|
||||
include = "include" if file.suffix == ".md" else "literalinclude"
|
||||
content += f":::{{admonition}} {file.relative_to(self.path)}\n"
|
||||
content += ":class: dropdown\n\n"
|
||||
content += f":::{{{include}}} {make_relative(file)}\n:::\n" # type: ignore[no-untyped-call]
|
||||
content += ":::\n\n"
|
||||
|
||||
return content
|
||||
|
||||
|
||||
def generate_examples():
|
||||
# Create the EXAMPLE_DOC_DIR if it doesn't exist
|
||||
if not EXAMPLE_DOC_DIR.exists():
|
||||
EXAMPLE_DOC_DIR.mkdir(parents=True)
|
||||
|
||||
# Create empty indices
|
||||
examples_index = Index(
|
||||
path=EXAMPLE_DOC_DIR / "examples_index.md",
|
||||
title="Examples",
|
||||
description=
|
||||
"A collection of examples demonstrating usage of FastVideo.\nAll documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.", # noqa: E501
|
||||
caption="Examples",
|
||||
maxdepth=2)
|
||||
# Category indices stored in reverse order because they are inserted into
|
||||
# examples_index.documents at index 0 in order
|
||||
category_indices = {
|
||||
"other":
|
||||
Index(
|
||||
path=EXAMPLE_DOC_DIR / "examples_other_index.md",
|
||||
title="Other",
|
||||
description=
|
||||
"Other examples that don't strongly fit into the online or offline serving categories.", # noqa: E501
|
||||
caption="Examples",
|
||||
),
|
||||
"online_serving":
|
||||
Index(
|
||||
path=EXAMPLE_DOC_DIR / "examples_online_serving_index.md",
|
||||
title="Online Serving",
|
||||
description=
|
||||
"Online serving examples demonstrate how to use FastVideo in an online setting, where the model is queried for predictions in real-time.", # noqa: E501
|
||||
caption="Examples",
|
||||
),
|
||||
"offline_inference":
|
||||
Index(
|
||||
path=EXAMPLE_DOC_DIR / "examples_offline_inference_index.md",
|
||||
title="Offline Inference",
|
||||
description=
|
||||
"Offline inference examples demonstrate how to use FastVideo in an offline setting, where the model is queried for predictions in batches. We recommend starting with <project:basic.md>.", # noqa: E501
|
||||
caption="Examples",
|
||||
),
|
||||
}
|
||||
|
||||
examples = []
|
||||
glob_patterns = ["*.py", "*.md", "*.sh"]
|
||||
# Find categorised examples
|
||||
for category in category_indices:
|
||||
category_dir = EXAMPLE_DIR / category
|
||||
globs = [category_dir.glob(pattern) for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
examples.append(Example(path, category))
|
||||
# Find examples in subdirectories
|
||||
for path in category_dir.glob("*/*.md"):
|
||||
examples.append(Example(path.parent, category))
|
||||
# Find uncategorised examples
|
||||
globs = [EXAMPLE_DIR.glob(pattern) for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
examples.append(Example(path))
|
||||
# Find examples in subdirectories
|
||||
for path in EXAMPLE_DIR.glob("*/*.md"):
|
||||
# Skip categorised examples
|
||||
if path.parent.name in category_indices:
|
||||
continue
|
||||
examples.append(Example(path.parent))
|
||||
|
||||
# Generate the example documentation
|
||||
for example in sorted(examples, key=lambda e: e.path.stem):
|
||||
doc_path = EXAMPLE_DOC_DIR / f"{example.path.stem}.md"
|
||||
with open(doc_path, "w+") as f:
|
||||
f.write(example.generate())
|
||||
# Add the example to the appropriate index
|
||||
assert example.category is not None
|
||||
index = category_indices.get(example.category, examples_index)
|
||||
index.documents.append(example.path.stem)
|
||||
|
||||
# Generate the index files
|
||||
for category_index in category_indices.values():
|
||||
if category_index.documents:
|
||||
examples_index.documents.insert(0, category_index.path.name)
|
||||
with open(category_index.path, "w+") as f:
|
||||
f.write(category_index.generate())
|
||||
|
||||
with open(examples_index.path, "w+") as f:
|
||||
f.write(examples_index.generate())
|
||||
@@ -1,10 +0,0 @@
|
||||
# Examples
|
||||
|
||||
A collection of examples demonstrating usage of FastVideo.
|
||||
All documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.
|
||||
|
||||
:::{toctree}
|
||||
:caption: Examples
|
||||
:maxdepth: 2
|
||||
|
||||
:::
|
||||
@@ -1,10 +0,0 @@
|
||||
(fastvideo-installation)=
|
||||
|
||||
# 🔧 Installation
|
||||
The code is tested on Python 3.10.0, CUDA 12.4 and H100.
|
||||
|
||||
```
|
||||
./env_setup.sh fastvideo
|
||||
```
|
||||
|
||||
To try Sliding Tile Attention (optional), please follow the instruction in [here](#sta-installation) to install STA.
|
||||
@@ -1,88 +0,0 @@
|
||||
# Welcome to FastVideo
|
||||
|
||||
:::{figure} ../../assets/logo.jpg
|
||||
:align: center
|
||||
:alt: FastVideo
|
||||
:class: no-scaled-link
|
||||
:width: 60%
|
||||
:::
|
||||
|
||||
:::{raw} html
|
||||
<p style="text-align:center">
|
||||
<strong>FastVideo is a lightweight framework for accelerating large video diffusion models.
|
||||
</strong>
|
||||
</p>
|
||||
|
||||
<p style="text-align:center">
|
||||
<script async defer src="https://buttons.github.io/buttons.js"></script>
|
||||
<a class="github-button" href="https://github.com/hao-ai-lab/FastVideo/" data-show-count="true" data-size="large" aria-label="Star">Star</a>
|
||||
<a class="github-button" href="https://github.com/hao-ai-lab/FastVideo/subscription" data-icon="octicon-eye" data-size="large" aria-label="Watch">Watch</a>
|
||||
<a class="github-button" href="https://github.com/hao-ai-lab/FastVideo/fork" data-icon="octicon-repo-forked" data-size="large" aria-label="Fork">Fork</a>
|
||||
</p>
|
||||
:::
|
||||
|
||||
FastVideo is a lightweight framework for accelerating large video diffusion models developed by the [Hao AI Lab](https://hao-ai-lab.github.io/).
|
||||
|
||||
<div style="text-align: center;">
|
||||
<video controls width="800">
|
||||
<source src="https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1" type="video/mp4">
|
||||
Your browser does not support the video tag.
|
||||
</video>
|
||||
</div>
|
||||
|
||||
FastVideo currently offers: (with more to come)
|
||||
|
||||
- [NEW!] [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- FastHunyuan and FastMochi: consistency distilled video diffusion models for 8x inference speedup.
|
||||
- First open distillation recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
|
||||
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
|
||||
- Scalable training with FSDP, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
- Memory efficient finetuning with LoRA, precomputed latent, and precomputed text embeddings.
|
||||
|
||||
Dev in progress and highly experimental.
|
||||
|
||||
## Documentation
|
||||
|
||||
% How to start using vLLM?
|
||||
|
||||
:::{toctree}
|
||||
:caption: Getting Started
|
||||
:maxdepth: 1
|
||||
|
||||
getting_started/installation
|
||||
getting_started/examples/examples_index
|
||||
:::
|
||||
|
||||
% What is STA Kernel?
|
||||
|
||||
:::{toctree}
|
||||
:caption: Sliding Tile Attention
|
||||
:maxdepth: 1
|
||||
|
||||
sliding_tile_attention/installation
|
||||
sliding_tile_attention/usage
|
||||
sliding_tile_attention/test
|
||||
sliding_tile_attention/demo
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Inference
|
||||
:maxdepth: 1
|
||||
|
||||
inference/stepvideo
|
||||
inference/hunyuanvideo
|
||||
inference/fasthunyuan
|
||||
inference/fastmochi
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Developer Guide
|
||||
:maxdepth: 1
|
||||
|
||||
developer_guide/overview
|
||||
:::
|
||||
|
||||
## Indices and tables
|
||||
|
||||
- {ref}`genindex`
|
||||
- {ref}`modindex`
|
||||
@@ -1,33 +0,0 @@
|
||||
(fasthunyuan)=
|
||||
|
||||
# FastHunyuan
|
||||
## Inference FastHunyuan on single RTX4090
|
||||
We now support NF4 and LLM-INT8 quantized inference using BitsAndBytes for FastHunyuan. With NF4 quantization, inference can be performed on a single RTX 4090 GPU, requiring just 20GB of VRAM.
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan-diffusers --local_dir=data/FastHunyuan-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan_hf_quantization.sh
|
||||
```
|
||||
|
||||
For more information about the VRAM requirements for BitsAndBytes quantization, please refer to the table below (timing measured on an H100 GPU):
|
||||
|
||||
| Configuration | Memory to Init Transformer | Peak Memory After Init Pipeline (Denoise) | Diffusion Time | End-to-End Time |
|
||||
|--------------------------------|----------------------------|--------------------------------------------|----------------|-----------------|
|
||||
| BF16 + Pipeline CPU Offload | 23.883G | 33.744G | 81s | 121.5s |
|
||||
| INT8 + Pipeline CPU Offload | 13.911G | 27.979G | 88s | 116.7s |
|
||||
| NF4 + Pipeline CPU Offload | 9.453G | 19.26G | 78s | 114.5s |
|
||||
|
||||
For improved quality in generated videos, we recommend using a GPU with 80GB of memory to run the BF16 model with the original Hunyuan pipeline. To execute the inference, use the following section:
|
||||
|
||||
## FastHunyuan
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan --local_dir=data/FastHunyuan --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan.sh
|
||||
```
|
||||
|
||||
You can also inference FastHunyuan in the [official Hunyuan github](https://github.com/Tencent/HunyuanVideo).
|
||||
@@ -1,9 +0,0 @@
|
||||
(fastmochi)=
|
||||
|
||||
# FastMochi
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastMochi-diffusers --local_dir=data/FastMochi-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_mochi_sp.sh
|
||||
@@ -1,18 +0,0 @@
|
||||
(hunyuanvideo)=
|
||||
|
||||
# HunyuanVideo
|
||||
## Inference HunyuanVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model
|
||||
```
|
||||
|
||||
We provide two examples in the following script to run inference with STA + [TeaCache](https://github.com/ali-vilab/TeaCache) and STA only.
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_hunyuan_STA.sh
|
||||
```
|
||||
|
||||
## Video Demos using STA + Teacache
|
||||
Visit our [demo website](https://fast-video.github.io/) to explore our complete collection of examples. We shorten a single video generation process from 945s to 317s on H100.
|
||||
@@ -1,16 +0,0 @@
|
||||
(stepvideo)=
|
||||
|
||||
# StepVideo
|
||||
## Inference StepVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
```
|
||||
python scripts/huggingface/download_hf.py --repo_id=stepfun-ai/stepvideo-t2v --local_dir=data/stepvideo-t2v --repo_type=model
|
||||
```
|
||||
|
||||
Use the following scripts to run inference for StepVideo. When using STA for inference, the generated videos will have dimensions of 204×768×768 (currently, this is the only supported shape).
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_stepvideo_STA.sh # Inference stepvideo with STA
|
||||
sh scripts/inference/inference_stepvideo.sh # Inference original stepvideo
|
||||
```
|
||||
@@ -1,11 +0,0 @@
|
||||
(sta-demo)=
|
||||
|
||||
# Demo
|
||||
There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
<div style="text-align: center;">
|
||||
<video controls width="800">
|
||||
<source src="https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747" type="video/mp4">
|
||||
Your browser does not support the video tag.
|
||||
</video>
|
||||
</div>
|
||||
@@ -1,25 +0,0 @@
|
||||
(sta-installation)=
|
||||
|
||||
# Installation
|
||||
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
|
||||
First, install C++20 for ThunderKittens:
|
||||
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
|
||||
Install STA:
|
||||
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
@@ -1,7 +0,0 @@
|
||||
(sta-test)=
|
||||
|
||||
# Test
|
||||
|
||||
```bash
|
||||
python test/test_sta.py
|
||||
```
|
||||
@@ -1,17 +0,0 @@
|
||||
(sta-usage)=
|
||||
|
||||
# Usage
|
||||
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
|
||||
```
|
||||
Executable
+14
@@ -0,0 +1,14 @@
|
||||
#!/bin/bash
|
||||
|
||||
# install torch
|
||||
pip install torch==2.5.0 torchvision --index-url https://download.pytorch.org/whl/cu124
|
||||
|
||||
# install FA2 and diffusers
|
||||
pip install packaging ninja && pip install flash-attn==2.7.0.post2 --no-build-isolation
|
||||
|
||||
pip install -r requirements-lint.txt
|
||||
|
||||
pip install -r requirements.txt
|
||||
|
||||
# install fastvideo
|
||||
pip install -e .
|
||||
@@ -1,17 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from .flash_attn import (DistributedAttention, LocalAttention)
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.attention.layer import DistributedAttention, LocalAttention
|
||||
from fastvideo.v1.attention.selector import get_attn_backend
|
||||
|
||||
__all__ = [
|
||||
"DistributedAttention",
|
||||
"LocalAttention",
|
||||
"AttentionBackend",
|
||||
"AttentionMetadata",
|
||||
"AttentionMetadataBuilder",
|
||||
# "AttentionState",
|
||||
"get_attn_backend",
|
||||
]
|
||||
__all__ = ["DistributedAttention", "LocalAttention"]
|
||||
@@ -1,245 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/attention/backends/abstract.py
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, fields
|
||||
from typing import (TYPE_CHECKING, Any, Dict, Generic, Optional, Protocol, Set,
|
||||
Type, TypeVar)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class AttentionBackend(ABC):
|
||||
"""Abstract class for attention backends."""
|
||||
# For some attention backends, we allocate an output tensor before
|
||||
# calling the custom op. When piecewise cudagraph is enabled, this
|
||||
# makes sure the output tensor is allocated inside the cudagraph.
|
||||
accept_output_buffer: bool = False
|
||||
|
||||
@staticmethod
|
||||
@abstractmethod
|
||||
def get_name() -> str:
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
@abstractmethod
|
||||
def get_impl_cls() -> Type["AttentionImpl"]:
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
@abstractmethod
|
||||
def get_metadata_cls() -> Type["AttentionMetadata"]:
|
||||
raise NotImplementedError
|
||||
|
||||
# @staticmethod
|
||||
# @abstractmethod
|
||||
# def get_state_cls() -> Type["AttentionState"]:
|
||||
# raise NotImplementedError
|
||||
|
||||
# @classmethod
|
||||
# def make_metadata(cls, *args, **kwargs) -> "AttentionMetadata":
|
||||
# return cls.get_metadata_cls()(*args, **kwargs)
|
||||
|
||||
@staticmethod
|
||||
@abstractmethod
|
||||
def get_builder_cls() -> Type["AttentionMetadataBuilder"]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@dataclass
|
||||
class AttentionMetadata:
|
||||
"""Attention metadata for prefill and decode batched together."""
|
||||
# Current step of diffusion process
|
||||
current_timestep: int
|
||||
|
||||
# @property
|
||||
# @abstractmethod
|
||||
# def inference_metadata(self) -> Optional["AttentionMetadata"]:
|
||||
# """Return the attention metadata that's required to run prefill
|
||||
# attention."""
|
||||
# pass
|
||||
|
||||
# @property
|
||||
# @abstractmethod
|
||||
# def training_metadata(self) -> Optional["AttentionMetadata"]:
|
||||
# """Return the attention metadata that's required to run decode
|
||||
# attention."""
|
||||
# pass
|
||||
|
||||
def asdict_zerocopy(self,
|
||||
skip_fields: Optional[Set[str]] = None
|
||||
) -> Dict[str, Any]:
|
||||
"""Similar to dataclasses.asdict, but avoids deepcopying."""
|
||||
if skip_fields is None:
|
||||
skip_fields = set()
|
||||
# Note that if we add dataclasses as fields, they will need
|
||||
# similar handling.
|
||||
return {
|
||||
field.name: getattr(self, field.name)
|
||||
for field in fields(self) if field.name not in skip_fields
|
||||
}
|
||||
|
||||
|
||||
T = TypeVar("T", bound=AttentionMetadata)
|
||||
|
||||
# class AttentionState(ABC, Generic[T]):
|
||||
# """Holds attention backend-specific objects reused during the
|
||||
# lifetime of the model runner."""
|
||||
|
||||
# @abstractmethod
|
||||
# def __init__(self, runner: "ModelRunnerBase"):
|
||||
# ...
|
||||
|
||||
# @abstractmethod
|
||||
# @contextmanager
|
||||
# def graph_capture(self, max_batch_size: int):
|
||||
# """Context manager used when capturing CUDA graphs."""
|
||||
# yield
|
||||
|
||||
# @abstractmethod
|
||||
# def graph_clone(self, batch_size: int) -> "AttentionState[T]":
|
||||
# """Clone attention state to save in CUDA graph metadata."""
|
||||
# ...
|
||||
|
||||
# @abstractmethod
|
||||
# def graph_capture_get_metadata_for_batch(
|
||||
# self,
|
||||
# batch_size: int,
|
||||
# is_encoder_decoder_model: bool = False) -> T:
|
||||
# """Get attention metadata for CUDA graph capture of batch_size."""
|
||||
# ...
|
||||
|
||||
# @abstractmethod
|
||||
# def get_graph_input_buffers(
|
||||
# self,
|
||||
# attn_metadata: T,
|
||||
# is_encoder_decoder_model: bool = False) -> Dict[str, Any]:
|
||||
# """Get attention-specific input buffers for CUDA graph capture."""
|
||||
# ...
|
||||
|
||||
# @abstractmethod
|
||||
# def prepare_graph_input_buffers(
|
||||
# self,
|
||||
# input_buffers: Dict[str, Any],
|
||||
# attn_metadata: T,
|
||||
# is_encoder_decoder_model: bool = False) -> None:
|
||||
# """In-place modify input buffers dict for CUDA graph replay."""
|
||||
# ...
|
||||
|
||||
# @abstractmethod
|
||||
# def begin_forward(self, model_input: "ModelRunnerInputBase") -> None:
|
||||
# """Prepare state for forward pass."""
|
||||
# ...
|
||||
|
||||
|
||||
class AttentionMetadataBuilder(ABC, Generic[T]):
|
||||
"""Abstract class for attention metadata builders."""
|
||||
|
||||
@abstractmethod
|
||||
def __init__(self) -> None:
|
||||
"""Create the builder, remember some configuration and parameters."""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def prepare(self) -> None:
|
||||
"""Prepare for one batch."""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def build(
|
||||
self,
|
||||
current_timestep: int,
|
||||
forward_batch: "ForwardBatch",
|
||||
inference_args: "InferenceArgs",
|
||||
) -> T:
|
||||
"""Build attention metadata with on-device tensors."""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class AttentionLayer(Protocol):
|
||||
|
||||
_k_scale: torch.Tensor
|
||||
_v_scale: torch.Tensor
|
||||
_k_scale_float: float
|
||||
_v_scale_float: float
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
...
|
||||
|
||||
|
||||
class AttentionImpl(ABC, Generic[T]):
|
||||
|
||||
@abstractmethod
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
softmax_scale: float,
|
||||
dropout_rate: float = 0.0,
|
||||
causal: bool = False,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
) -> None:
|
||||
raise NotImplementedError
|
||||
|
||||
def preprocess_qkv(self, qkv: torch.Tensor,
|
||||
attn_metadata: T) -> torch.Tensor:
|
||||
"""Preprocess QKV tensor before performing attention operation.
|
||||
|
||||
Default implementation returns the tensor unchanged.
|
||||
Subclasses can override this to implement custom preprocessing
|
||||
like reshaping, tiling, scaling, or other transformations.
|
||||
|
||||
Called AFTER all_to_all for distributed attention
|
||||
|
||||
Args:
|
||||
qkv: The query-key-value tensor
|
||||
attn_metadata: Metadata for the attention operation
|
||||
|
||||
Returns:
|
||||
Processed QKV tensor
|
||||
"""
|
||||
return qkv
|
||||
|
||||
def postprocess_output(
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
attn_metadata: T,
|
||||
) -> torch.Tensor:
|
||||
"""Postprocess the output tensor after the attention operation.
|
||||
|
||||
Default implementation returns the tensor unchanged.
|
||||
Subclasses can override this to implement custom postprocessing
|
||||
like untiling, scaling, or other transformations.
|
||||
|
||||
Called BEFORE all_to_all for distributed attention
|
||||
|
||||
Args:
|
||||
output: The output tensor from the attention operation
|
||||
attn_metadata: Metadata for the attention operation
|
||||
|
||||
Returns:
|
||||
Postprocessed output tensor
|
||||
"""
|
||||
|
||||
return output
|
||||
|
||||
@abstractmethod
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: T,
|
||||
) -> torch.Tensor:
|
||||
raise NotImplementedError
|
||||
@@ -1,70 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
from flash_attn import flash_attn_func
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class FlashAttentionBackend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> List[int]:
|
||||
return [32, 64, 96, 128, 160, 192, 224, 256]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "FLASH_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> Type["FlashAttentionImpl"]:
|
||||
return FlashAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> Type["AttentionMetadata"]:
|
||||
raise NotImplementedError
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> Type["AttentionMetadataBuilder"]:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class FlashAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
dropout_rate: float,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
) -> None:
|
||||
self.dropout_rate = dropout_rate
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
):
|
||||
output = flash_attn_func(query,
|
||||
key,
|
||||
value,
|
||||
dropout_p=self.dropout_rate,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal)
|
||||
return output
|
||||
@@ -1,72 +0,0 @@
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (
|
||||
AttentionBackend) # FlashAttentionMetadata,
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionImpl,
|
||||
AttentionMetadata)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class SDPABackend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> List[int]:
|
||||
return [32, 64, 96, 128, 160, 192, 224, 256]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SDPA"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> Type["SDPAImpl"]:
|
||||
return SDPAImpl
|
||||
|
||||
# @staticmethod
|
||||
# def get_metadata_cls() -> Type["AttentionMetadata"]:
|
||||
# return FlashAttentionMetadata
|
||||
|
||||
|
||||
class SDPAImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
dropout_rate: float,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
) -> None:
|
||||
self.dropout_rate = dropout_rate
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
|
||||
def forward(
|
||||
self,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
value: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
# transpose to bs, heads, seq_len, head_dim
|
||||
query = query.transpose(1, 2)
|
||||
key = key.transpose(1, 2)
|
||||
value = value.transpose(1, 2)
|
||||
attn_kwargs = {
|
||||
"attn_mask": None,
|
||||
"dropout_p": self.dropout_rate,
|
||||
"is_causal": self.causal,
|
||||
"scale": self.softmax_scale
|
||||
}
|
||||
if query.shape[1] != key.shape[1]:
|
||||
attn_kwargs["enable_gqa"] = True
|
||||
output = torch.nn.functional.scaled_dot_product_attention(
|
||||
query, key, value, **attn_kwargs)
|
||||
output = output.transpose(1, 2)
|
||||
return output
|
||||
@@ -1,195 +0,0 @@
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from st_attn import sliding_tile_attention
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.distributed import get_sp_group
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
# TODO(will-refactor): move this to a utils file
|
||||
def dict_to_3d_list(mask_strategy,
|
||||
t_max=50,
|
||||
l_max=60,
|
||||
h_max=24) -> List[List[List[Optional[torch.Tensor]]]]:
|
||||
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, layer, h = map(int, key.split('_'))
|
||||
result[t][layer][h] = value
|
||||
return result
|
||||
|
||||
|
||||
class SlidingTileAttentionBackend(AttentionBackend):
|
||||
|
||||
accept_output_buffer: bool = True
|
||||
|
||||
@staticmethod
|
||||
def get_supported_head_sizes() -> List[int]:
|
||||
# TODO(will-refactor): check this
|
||||
return [32, 64, 96, 128, 160, 192, 224, 256]
|
||||
|
||||
@staticmethod
|
||||
def get_name() -> str:
|
||||
return "SLIDING_TILE_ATTN"
|
||||
|
||||
@staticmethod
|
||||
def get_impl_cls() -> Type["SlidingTileAttentionImpl"]:
|
||||
return SlidingTileAttentionImpl
|
||||
|
||||
@staticmethod
|
||||
def get_metadata_cls() -> Type["SlidingTileAttentionMetadata"]:
|
||||
return SlidingTileAttentionMetadata
|
||||
|
||||
@staticmethod
|
||||
def get_builder_cls() -> Type["SlidingTileAttentionMetadataBuilder"]:
|
||||
return SlidingTileAttentionMetadataBuilder
|
||||
|
||||
|
||||
@dataclass
|
||||
class SlidingTileAttentionMetadata(AttentionMetadata):
|
||||
text_length: int
|
||||
|
||||
|
||||
class SlidingTileAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
|
||||
def __init__(self):
|
||||
pass
|
||||
|
||||
def prepare(self):
|
||||
pass
|
||||
|
||||
def build(
|
||||
self,
|
||||
current_timestep: int,
|
||||
forward_batch: ForwardBatch,
|
||||
inference_args: InferenceArgs,
|
||||
) -> SlidingTileAttentionMetadata:
|
||||
|
||||
return SlidingTileAttentionMetadata(
|
||||
current_timestep=current_timestep,
|
||||
text_length=forward_batch.attention_mask.sum(),
|
||||
)
|
||||
|
||||
|
||||
class SlidingTileAttentionImpl(AttentionImpl):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
dropout_rate: float,
|
||||
causal: bool,
|
||||
softmax_scale: float,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
) -> None:
|
||||
# TODO(will-refactor): for now this is the mask strategy, but maybe we should
|
||||
# have a more general config for STA?
|
||||
config_file = envs.FASTVIDEO_ATTENTION_CONFIG
|
||||
if config_file is None:
|
||||
raise ValueError("FASTVIDEO_ATTENTION_CONFIG is not set")
|
||||
|
||||
with open(config_file) as f:
|
||||
mask_strategy = json.load(f)
|
||||
|
||||
mask_strategy = dict_to_3d_list(mask_strategy)
|
||||
|
||||
self.mask_strategy = mask_strategy
|
||||
sp_group = get_sp_group()
|
||||
self.sp_size = sp_group.world_size
|
||||
|
||||
def tile(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = rearrange(x,
|
||||
"b (sp t h w) head d -> b (t sp h w) head d",
|
||||
sp=self.sp_size,
|
||||
t=30 // self.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(self, x: torch.Tensor) -> torch.Tensor:
|
||||
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=self.sp_size,
|
||||
t=30 // self.sp_size,
|
||||
h=48,
|
||||
w=80)
|
||||
|
||||
def preprocess_qkv(
|
||||
self,
|
||||
qkv: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
return self.tile(qkv)
|
||||
|
||||
def postprocess_output(
|
||||
self,
|
||||
output: torch.Tensor,
|
||||
attn_metadata: SlidingTileAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
return self.untile(output)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
attn_metadata: SlidingTileAttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
|
||||
assert self.mask_strategy is not None, "mask_strategy cannot be None for SlidingTileAttention"
|
||||
assert self.mask_strategy[
|
||||
0] is not None, "mask_strategy[0] cannot be None for SlidingTileAttention"
|
||||
|
||||
text_length = attn_metadata.text_length
|
||||
|
||||
query = q.transpose(1, 2)
|
||||
key = k.transpose(1, 2)
|
||||
value = v.transpose(1, 2)
|
||||
|
||||
head_num = query.size(1)
|
||||
sp_group = get_sp_group()
|
||||
current_rank = sp_group.rank_in_group
|
||||
start_head = current_rank * head_num
|
||||
windows = [
|
||||
self.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)
|
||||
|
||||
hidden_states = hidden_states.transpose(1, 2)
|
||||
|
||||
return hidden_states
|
||||
@@ -0,0 +1,138 @@
|
||||
from itertools import accumulate
|
||||
from typing import List, Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from fastvideo.v1.distributed.communication_op import sequence_model_parallel_all_to_all_4D, sequence_model_parallel_all_gather
|
||||
from fastvideo.v1.distributed.parallel_state import get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size
|
||||
from flash_attn import flash_attn_func, flash_attn_varlen_func
|
||||
|
||||
|
||||
class DistributedAttention(nn.Module):
|
||||
"""Distributed attention module that supports sequence parallelism.
|
||||
|
||||
This class implements a minimal attention operation with support for distributed
|
||||
processing across multiple GPUs using sequence parallelism. The implementation assumes
|
||||
batch_size=1 and no padding tokens for simplicity.
|
||||
|
||||
The sequence parallelism strategy follows the Ulysses paper (https://arxiv.org/abs/2309.14509),
|
||||
which proposes redistributing attention heads across sequence dimension to enable efficient
|
||||
parallel processing of long sequences.
|
||||
|
||||
Args:
|
||||
dropout_rate (float, optional): Dropout probability. Defaults to 0.0.
|
||||
causal (bool, optional): Whether to use causal attention. Defaults to False.
|
||||
softmax_scale (float, optional): Custom scaling factor for attention scores.
|
||||
If None, uses 1/sqrt(head_dim). Defaults to None.
|
||||
"""
|
||||
def __init__(
|
||||
self,
|
||||
dropout_rate: float = 0.0,
|
||||
causal: bool = False,
|
||||
softmax_scale: Optional[float] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.dropout_rate = dropout_rate
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
|
||||
def forward(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
replicated_q: Optional[torch.Tensor] = None,
|
||||
replicated_k: Optional[torch.Tensor] = None,
|
||||
replicated_v: Optional[torch.Tensor] = None,
|
||||
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||
"""Forward pass for distributed attention.
|
||||
|
||||
Args:
|
||||
q (torch.Tensor): Query tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
k (torch.Tensor): Key tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
v (torch.Tensor): Value tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
replicated_q (Optional[torch.Tensor]): Replicated query tensor, typically for text tokens
|
||||
replicated_k (Optional[torch.Tensor]): Replicated key tensor
|
||||
replicated_v (Optional[torch.Tensor]): Replicated value tensor
|
||||
|
||||
Returns:
|
||||
Tuple[torch.Tensor, Optional[torch.Tensor]]: A tuple containing:
|
||||
- o (torch.Tensor): Output tensor after attention for the main sequence
|
||||
- replicated_o (Optional[torch.Tensor]): Output tensor for replicated tokens, if provided
|
||||
"""
|
||||
# Check input shapes
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Expected 4D tensors"
|
||||
# assert bs = 1
|
||||
assert q.shape[0] == 1, "Batch size must be 1, and there should be no padding tokens"
|
||||
batch_size, seq_len, num_heads, head_dim = q.shape
|
||||
local_rank = get_sequence_model_parallel_rank()
|
||||
world_size = get_sequence_model_parallel_world_size()
|
||||
|
||||
# Stack QKV
|
||||
qkv = torch.cat([q, k, v], dim=0) # [3, seq_len, num_heads, head_dim]
|
||||
|
||||
# Redistribute heads across sequence dimension
|
||||
qkv = sequence_model_parallel_all_to_all_4D(qkv, scatter_dim=2, gather_dim=1)
|
||||
|
||||
# Concatenate with replicated QKV if provided
|
||||
if replicated_q is not None:
|
||||
assert replicated_k is not None and replicated_v is not None
|
||||
replicated_qkv = torch.cat([replicated_q, replicated_k, replicated_v], dim=0) # [3, seq_len, num_heads, head_dim]
|
||||
heads_per_rank = num_heads // world_size
|
||||
replicated_qkv = replicated_qkv[:, :, local_rank * heads_per_rank:(local_rank + 1) * heads_per_rank]
|
||||
qkv = torch.cat([qkv, replicated_qkv], dim=1)
|
||||
|
||||
q, k, v = qkv.chunk(3, dim=0)
|
||||
# Apply flash attention
|
||||
output = flash_attn_func(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
dropout_p=self.dropout_rate,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal
|
||||
)
|
||||
# Redistribute back if using sequence parallelism
|
||||
replicated_output = None
|
||||
if replicated_q is not None:
|
||||
replicated_output = output[:, seq_len*world_size:]
|
||||
output = output[:, :seq_len*world_size]
|
||||
# TODO: make this asynchronous
|
||||
replicated_output = sequence_model_parallel_all_gather(replicated_output, dim=2)
|
||||
output = sequence_model_parallel_all_to_all_4D(output, scatter_dim=1, gather_dim=2)
|
||||
return output, replicated_output
|
||||
|
||||
|
||||
class LocalAttention(nn.Module):
|
||||
def __init__(self, dropout_rate: float = 0.0, causal: bool = False, softmax_scale: Optional[float] = None):
|
||||
super().__init__()
|
||||
self.dropout_rate = dropout_rate
|
||||
self.causal = causal
|
||||
self.softmax_scale = softmax_scale
|
||||
|
||||
def forward(self, q, k, v):
|
||||
"""
|
||||
Apply local attention between query, key and value tensors.
|
||||
|
||||
Args:
|
||||
q (torch.Tensor): Query tensor of shape [batch_size, seq_len, num_heads, head_dim]
|
||||
k (torch.Tensor): Key tensor of shape [batch_size, seq_len, num_heads, head_dim]
|
||||
v (torch.Tensor): Value tensor of shape [batch_size, seq_len, num_heads, head_dim]
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor after local attention
|
||||
"""
|
||||
# Check input shapes
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "Expected 4D tensors"
|
||||
|
||||
# Apply flash attention
|
||||
output = flash_attn_func(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
dropout_p=self.dropout_rate,
|
||||
softmax_scale=self.softmax_scale,
|
||||
causal=self.causal
|
||||
)
|
||||
|
||||
return output
|
||||
@@ -1,201 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.v1.attention.selector import (backend_name_to_enum,
|
||||
get_attn_backend)
|
||||
from fastvideo.v1.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather, sequence_model_parallel_all_to_all_4D)
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
|
||||
|
||||
|
||||
class DistributedAttention(nn.Module):
|
||||
"""Distributed attention layer.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
dropout_rate: float = 0.0,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
# self.dropout_rate = dropout_rate
|
||||
# self.causal = causal
|
||||
if softmax_scale is None:
|
||||
self.softmax_scale = head_size**-0.5
|
||||
else:
|
||||
self.softmax_scale = softmax_scale
|
||||
|
||||
if num_kv_heads is None:
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = torch.get_default_dtype()
|
||||
attn_backend = get_attn_backend(head_size, dtype, distributed=True)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
dropout_rate=dropout_rate,
|
||||
causal=causal,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
**extra_impl_args)
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
self.num_kv_heads = num_kv_heads
|
||||
self.backend = backend_name_to_enum(attn_backend.get_name())
|
||||
self.dtype = dtype
|
||||
|
||||
def forward(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
replicated_q: Optional[torch.Tensor] = None,
|
||||
replicated_k: Optional[torch.Tensor] = None,
|
||||
replicated_v: Optional[torch.Tensor] = None,
|
||||
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
|
||||
"""Forward pass for distributed attention.
|
||||
|
||||
Args:
|
||||
q (torch.Tensor): Query tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
k (torch.Tensor): Key tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
v (torch.Tensor): Value tensor [batch_size, seq_len, num_heads, head_dim]
|
||||
replicated_q (Optional[torch.Tensor]): Replicated query tensor, typically for text tokens
|
||||
replicated_k (Optional[torch.Tensor]): Replicated key tensor
|
||||
replicated_v (Optional[torch.Tensor]): Replicated value tensor
|
||||
|
||||
Returns:
|
||||
Tuple[torch.Tensor, Optional[torch.Tensor]]: A tuple containing:
|
||||
- o (torch.Tensor): Output tensor after attention for the main sequence
|
||||
- replicated_o (Optional[torch.Tensor]): Output tensor for replicated tokens, if provided
|
||||
"""
|
||||
# Check input shapes
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim(
|
||||
) == 4, "Expected 4D tensors"
|
||||
# assert bs = 1
|
||||
assert q.shape[
|
||||
0] == 1, "Batch size must be 1, and there should be no padding tokens"
|
||||
batch_size, seq_len, num_heads, head_dim = q.shape
|
||||
local_rank = get_sequence_model_parallel_rank()
|
||||
world_size = get_sequence_model_parallel_world_size()
|
||||
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
|
||||
# Stack QKV
|
||||
qkv = torch.cat([q, k, v], dim=0) # [3, seq_len, num_heads, head_dim]
|
||||
|
||||
# Redistribute heads across sequence dimension
|
||||
qkv = sequence_model_parallel_all_to_all_4D(qkv,
|
||||
scatter_dim=2,
|
||||
gather_dim=1)
|
||||
|
||||
# Apply backend-specific preprocess_qkv
|
||||
qkv = self.impl.preprocess_qkv(qkv, ctx_attn_metadata)
|
||||
|
||||
# Concatenate with replicated QKV if provided
|
||||
if replicated_q is not None:
|
||||
assert replicated_k is not None and replicated_v is not None
|
||||
replicated_qkv = torch.cat(
|
||||
[replicated_q, replicated_k, replicated_v],
|
||||
dim=0) # [3, seq_len, num_heads, head_dim]
|
||||
heads_per_rank = num_heads // world_size
|
||||
replicated_qkv = replicated_qkv[:, :, local_rank *
|
||||
heads_per_rank:(local_rank + 1) *
|
||||
heads_per_rank]
|
||||
qkv = torch.cat([qkv, replicated_qkv], dim=1)
|
||||
|
||||
q, k, v = qkv.chunk(3, dim=0)
|
||||
|
||||
output = self.impl.forward(q, k, v, ctx_attn_metadata)
|
||||
|
||||
# Redistribute back if using sequence parallelism
|
||||
replicated_output = None
|
||||
if replicated_q is not None:
|
||||
replicated_output = output[:, seq_len * world_size:]
|
||||
output = output[:, :seq_len * world_size]
|
||||
# TODO: make this asynchronous
|
||||
replicated_output = sequence_model_parallel_all_gather(
|
||||
replicated_output, dim=2)
|
||||
|
||||
# Apply backend-specific postprocess_output
|
||||
output = self.impl.postprocess_output(output, ctx_attn_metadata)
|
||||
|
||||
output = sequence_model_parallel_all_to_all_4D(output,
|
||||
scatter_dim=1,
|
||||
gather_dim=2)
|
||||
return output, replicated_output
|
||||
|
||||
|
||||
class LocalAttention(nn.Module):
|
||||
"""Attention layer.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
num_heads: int,
|
||||
head_size: int,
|
||||
num_kv_heads: Optional[int] = None,
|
||||
dropout_rate: float = 0.0,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
# self.dropout_rate = dropout_rate
|
||||
# self.causal = causal
|
||||
if softmax_scale is None:
|
||||
self.softmax_scale = head_size**-0.5
|
||||
else:
|
||||
self.softmax_scale = softmax_scale
|
||||
if num_kv_heads is None:
|
||||
num_kv_heads = num_heads
|
||||
|
||||
dtype = torch.get_default_dtype()
|
||||
attn_backend = get_attn_backend(head_size, dtype, distributed=False)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
dropout_rate=dropout_rate,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
causal=causal,
|
||||
**extra_impl_args)
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
self.num_kv_heads = num_kv_heads
|
||||
self.backend = backend_name_to_enum(attn_backend.get_name())
|
||||
self.dtype = dtype
|
||||
|
||||
def forward(
|
||||
self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Apply local attention between query, key and value tensors.
|
||||
|
||||
Args:
|
||||
q (torch.Tensor): Query tensor of shape [batch_size, seq_len, num_heads, head_dim]
|
||||
k (torch.Tensor): Key tensor of shape [batch_size, seq_len, num_heads, head_dim]
|
||||
v (torch.Tensor): Value tensor of shape [batch_size, seq_len, num_heads, head_dim]
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Output tensor after local attention
|
||||
"""
|
||||
# Check input shapes
|
||||
assert q.dim() == 4 and k.dim() == 4 and v.dim(
|
||||
) == 4, "Expected 4D tensors"
|
||||
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
|
||||
output = self.impl.forward(q, k, v, ctx_attn_metadata)
|
||||
return output
|
||||
@@ -1,157 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/attention/selector.py
|
||||
|
||||
import os
|
||||
from contextlib import contextmanager
|
||||
from functools import cache
|
||||
from typing import Generator, Optional, Type, cast
|
||||
|
||||
import torch
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.attention.backends.abstract import AttentionBackend
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.platforms import _Backend, current_platform
|
||||
from fastvideo.v1.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
|
||||
"""
|
||||
Convert a string backend name to a _Backend enum value.
|
||||
|
||||
Returns:
|
||||
* _Backend: enum value if backend_name is a valid in-tree type
|
||||
* None: otherwise it's an invalid in-tree type or an out-of-tree platform is
|
||||
loaded.
|
||||
"""
|
||||
assert backend_name is not None
|
||||
return _Backend[backend_name] if backend_name in _Backend.__members__ else \
|
||||
None
|
||||
|
||||
|
||||
def get_env_variable_attn_backend() -> Optional[_Backend]:
|
||||
'''
|
||||
Get the backend override specified by the FastVideo attention
|
||||
backend environment variable, if one is specified.
|
||||
|
||||
Returns:
|
||||
|
||||
* _Backend enum value if an override is specified
|
||||
* None otherwise
|
||||
'''
|
||||
backend_name = os.environ.get(STR_BACKEND_ENV_VAR)
|
||||
return (None
|
||||
if backend_name is None else backend_name_to_enum(backend_name))
|
||||
|
||||
|
||||
# Global state allows a particular choice of backend
|
||||
# to be forced, overriding the logic which auto-selects
|
||||
# a backend based on system & workload configuration
|
||||
# (default behavior if this variable is None)
|
||||
#
|
||||
# THIS SELECTION TAKES PRECEDENCE OVER THE
|
||||
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
|
||||
forced_attn_backend: Optional[_Backend] = None
|
||||
|
||||
|
||||
def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
|
||||
'''
|
||||
Force all attention operations to use a specified backend.
|
||||
|
||||
Passing `None` for the argument re-enables automatic
|
||||
backend selection.,
|
||||
|
||||
Arguments:
|
||||
|
||||
* attn_backend: backend selection (None to revert to auto)
|
||||
'''
|
||||
global forced_attn_backend
|
||||
forced_attn_backend = attn_backend
|
||||
|
||||
|
||||
def get_global_forced_attn_backend() -> Optional[_Backend]:
|
||||
'''
|
||||
Get the currently-forced choice of attention backend,
|
||||
or None if auto-selection is currently enabled.
|
||||
'''
|
||||
return forced_attn_backend
|
||||
|
||||
|
||||
def get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
distributed: bool,
|
||||
) -> Type[AttentionBackend]:
|
||||
"""Selects which attention backend to use and lazily imports it."""
|
||||
# Accessing envs.* behind an @lru_cache decorator can cause the wrong
|
||||
# value to be returned from the cache if the value changes between calls.
|
||||
return _cached_get_attn_backend(
|
||||
head_size=head_size,
|
||||
dtype=dtype,
|
||||
distributed=distributed,
|
||||
)
|
||||
|
||||
|
||||
@cache
|
||||
def _cached_get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
distributed: bool,
|
||||
) -> Type[AttentionBackend]:
|
||||
# Check whether a particular choice of backend was
|
||||
# previously forced.
|
||||
#
|
||||
# THIS SELECTION OVERRIDES THE FASTVIDEO_ATTENTION_BACKEND
|
||||
# ENVIRONMENT VARIABLE.
|
||||
selected_backend = None
|
||||
backend_by_global_setting: Optional[_Backend] = (
|
||||
get_global_forced_attn_backend())
|
||||
if backend_by_global_setting is not None:
|
||||
selected_backend = backend_by_global_setting
|
||||
else:
|
||||
# Check the environment variable and override if specified
|
||||
backend_by_env_var: Optional[str] = envs.FASTVIDEO_ATTENTION_BACKEND
|
||||
if backend_by_env_var is not None:
|
||||
selected_backend = backend_name_to_enum(backend_by_env_var)
|
||||
|
||||
# get device-specific attn_backend
|
||||
attention_cls = current_platform.get_attn_backend_cls(
|
||||
selected_backend, head_size, dtype, distributed)
|
||||
if not attention_cls:
|
||||
raise ValueError(
|
||||
f"Invalid attention backend for {current_platform.device_name}")
|
||||
return cast(Type[AttentionBackend], resolve_obj_by_qualname(attention_cls))
|
||||
|
||||
|
||||
@contextmanager
|
||||
def global_force_attn_backend_context_manager(
|
||||
attn_backend: _Backend) -> Generator[None, None, None]:
|
||||
'''
|
||||
Globally force a FastVideo attention backend override within a
|
||||
context manager, reverting the global attention backend
|
||||
override to its prior state upon exiting the context
|
||||
manager.
|
||||
|
||||
Arguments:
|
||||
|
||||
* attn_backend: attention backend to force
|
||||
|
||||
Returns:
|
||||
|
||||
* Generator
|
||||
'''
|
||||
|
||||
# Save the current state of the global backend override (if any)
|
||||
original_value = get_global_forced_attn_backend()
|
||||
|
||||
# Globally force the new backend override
|
||||
global_force_attn_backend(attn_backend)
|
||||
|
||||
# Yield control back to the enclosed code block
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
# Revert the original global backend override, if any
|
||||
global_force_attn_backend(original_value)
|
||||
@@ -1,16 +0,0 @@
|
||||
num_gpus: 4
|
||||
model_path: FastVideo/FastHunyuan-diffusers
|
||||
master_port: 29503
|
||||
sp_size: 4
|
||||
tp_size: 4
|
||||
height: 720
|
||||
width: 1280
|
||||
num_frames: 125
|
||||
num_inference_steps: 6
|
||||
guidance_scale: 1
|
||||
embedded_cfg_scale: 6
|
||||
flow_shift: 17
|
||||
prompt_path: ./assets/prompt.txt
|
||||
seed: 1024
|
||||
output_path: outputs_video/
|
||||
vae-sp: True
|
||||
@@ -1,5 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.v1.distributed.communication_op import *
|
||||
from fastvideo.v1.distributed.parallel_state import *
|
||||
from fastvideo.v1.distributed.utils import *
|
||||
from .communication_op import *
|
||||
from .parallel_state import *
|
||||
from .utils import *
|
||||
|
||||
@@ -1,10 +1,12 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/distributed/communication_op.py
|
||||
|
||||
from typing import Any, Dict, Optional, Union
|
||||
|
||||
import torch
|
||||
import torch.distributed
|
||||
|
||||
from fastvideo.v1.distributed.parallel_state import get_sp_group, get_tp_group
|
||||
from .parallel_state import get_tp_group, get_sp_group
|
||||
|
||||
|
||||
def tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
|
||||
@@ -18,6 +20,21 @@ def tensor_model_parallel_all_gather(input_: torch.Tensor,
|
||||
return get_tp_group().all_gather(input_, dim)
|
||||
|
||||
|
||||
def tensor_model_parallel_gather(input_: torch.Tensor,
|
||||
dst: int = 0,
|
||||
dim: int = -1) -> Optional[torch.Tensor]:
|
||||
"""Gather the input tensor across model parallel group."""
|
||||
return get_tp_group().gather(input_, dst, dim)
|
||||
|
||||
|
||||
def broadcast_tensor_dict(tensor_dict: Optional[Dict[Any, Union[torch.Tensor,
|
||||
Any]]] = None,
|
||||
src: int = 0):
|
||||
if not torch.distributed.is_initialized():
|
||||
return tensor_dict
|
||||
return get_tp_group().broadcast_tensor_dict(tensor_dict, src)
|
||||
|
||||
|
||||
# TODO: remove model, make it sequence_parallel
|
||||
def sequence_model_parallel_all_to_all_4D(input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
@@ -27,6 +44,8 @@ def sequence_model_parallel_all_to_all_4D(input_: torch.Tensor,
|
||||
|
||||
|
||||
def sequence_model_parallel_all_gather(input_: torch.Tensor,
|
||||
dim: int = -1) -> torch.Tensor:
|
||||
dim: int = -1) -> torch.Tensor:
|
||||
"""All-gather the input tensor across model parallel group."""
|
||||
return get_sp_group().all_gather(input_, dim)
|
||||
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ from typing import Optional
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import ProcessGroup
|
||||
|
||||
from einops import rearrange
|
||||
|
||||
class DeviceCommunicatorBase:
|
||||
"""
|
||||
@@ -94,11 +94,10 @@ class DeviceCommunicatorBase:
|
||||
else:
|
||||
output_tensor = None
|
||||
return output_tensor
|
||||
|
||||
def all_to_all_4D(self,
|
||||
input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1) -> torch.Tensor:
|
||||
def all_to_all_4D(self,
|
||||
input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1) -> torch.Tensor:
|
||||
"""Specialized all-to-all operation for 4D tensors (e.g., for QKV matrices).
|
||||
|
||||
Args:
|
||||
@@ -112,64 +111,55 @@ class DeviceCommunicatorBase:
|
||||
# Bypass the function if we are using only 1 GPU.
|
||||
if self.world_size == 1:
|
||||
return input_
|
||||
|
||||
assert input_.dim(
|
||||
) == 4, f"input must be 4D tensor, got {input_.dim()} and shape {input_.shape}"
|
||||
|
||||
|
||||
assert input_.dim() == 4, f"input must be 4D tensor, got {input_.dim()} and shape {input_.shape}"
|
||||
|
||||
if scatter_dim == 2 and gather_dim == 1:
|
||||
# input: (bs, seqlen/P, hc, hs) output: (bs, seqlen, hc/P, hs)
|
||||
bs, shard_seqlen, hc, hs = input_.shape
|
||||
seqlen = shard_seqlen * self.world_size
|
||||
shard_hc = hc // self.world_size
|
||||
|
||||
|
||||
# Reshape and transpose for scattering
|
||||
input_t = (input_.reshape(bs, shard_seqlen, self.world_size,
|
||||
shard_hc, hs).transpose(0,
|
||||
2).contiguous())
|
||||
|
||||
input_t = (input_.reshape(bs, shard_seqlen, self.world_size, shard_hc, hs).transpose(0, 2).contiguous())
|
||||
|
||||
output = torch.empty_like(input_t)
|
||||
|
||||
|
||||
torch.distributed.all_to_all_single(output,
|
||||
input_t,
|
||||
group=self.device_group)
|
||||
torch.distributed.all_to_all_single(output, input_t, group=self.device_group)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
|
||||
# Reshape and transpose back
|
||||
output = output.reshape(seqlen, bs, shard_hc,
|
||||
hs).transpose(0, 1).contiguous().reshape(
|
||||
bs, seqlen, shard_hc, hs)
|
||||
|
||||
output = output.reshape(seqlen, bs, shard_hc, hs).transpose(0, 1).contiguous().reshape(bs, seqlen, shard_hc, hs)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
elif scatter_dim == 1 and gather_dim == 2:
|
||||
# input: (bs, seqlen, hc/P, hs) output: (bs, seqlen/P, hc, hs)
|
||||
bs, seqlen, shard_hc, hs = input_.shape
|
||||
hc = shard_hc * self.world_size
|
||||
shard_seqlen = seqlen // self.world_size
|
||||
|
||||
|
||||
# Reshape and transpose for scattering
|
||||
input_t = (input_.reshape(bs, self.world_size, shard_seqlen,
|
||||
shard_hc, hs).transpose(0, 3).transpose(
|
||||
0, 1).contiguous().reshape(
|
||||
self.world_size, shard_hc,
|
||||
shard_seqlen, bs, hs))
|
||||
input_t = (input_.reshape(bs, self.world_size, shard_seqlen, shard_hc,
|
||||
hs).transpose(0,
|
||||
3).transpose(0,
|
||||
1).contiguous().reshape(self.world_size, shard_hc,
|
||||
shard_seqlen, bs, hs))
|
||||
output = torch.empty_like(input_t)
|
||||
|
||||
|
||||
torch.distributed.all_to_all_single(output,
|
||||
input_t,
|
||||
group=self.device_group)
|
||||
torch.distributed.all_to_all_single(output, input_t, group=self.device_group)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
|
||||
# Reshape and transpose back
|
||||
output = output.reshape(hc, shard_seqlen, bs,
|
||||
hs).transpose(0, 2).contiguous().reshape(
|
||||
bs, shard_seqlen, hc, hs)
|
||||
|
||||
output = output.reshape(hc, shard_seqlen, bs, hs).transpose(0, 2).contiguous().reshape(bs, shard_seqlen, hc, hs)
|
||||
|
||||
return output
|
||||
else:
|
||||
raise RuntimeError(
|
||||
"scatter_dim must be 1 or 2 and gather_dim must be 1 or 2")
|
||||
|
||||
raise RuntimeError("scatter_dim must be 1 or 2 and gather_dim must be 1 or 2")
|
||||
|
||||
|
||||
def send(self, tensor: torch.Tensor, dst: Optional[int] = None) -> None:
|
||||
"""Sends a tensor to the destination rank in a non-blocking way"""
|
||||
"""NOTE: `dst` is the local rank of the destination rank."""
|
||||
|
||||
@@ -6,8 +6,7 @@ from typing import Optional
|
||||
import torch
|
||||
from torch.distributed import ProcessGroup
|
||||
|
||||
from fastvideo.v1.distributed.device_communicators.base_device_communicator import (
|
||||
DeviceCommunicatorBase)
|
||||
from .base_device_communicator import DeviceCommunicatorBase
|
||||
|
||||
|
||||
class CudaCommunicator(DeviceCommunicatorBase):
|
||||
@@ -18,18 +17,50 @@ class CudaCommunicator(DeviceCommunicatorBase):
|
||||
device_group: Optional[ProcessGroup] = None,
|
||||
unique_name: str = ""):
|
||||
super().__init__(cpu_group, device, device_group, unique_name)
|
||||
if "pp" in unique_name:
|
||||
# pipeline parallel does not need custom allreduce
|
||||
use_custom_allreduce = False
|
||||
else:
|
||||
# from vllm.distributed.parallel_state import (
|
||||
# _ENABLE_CUSTOM_ALL_REDUCE)
|
||||
# TODO(will): bring in the custom allreduce from vLLM
|
||||
use_custom_allreduce = False
|
||||
use_pynccl = True
|
||||
|
||||
self.use_pynccl = use_pynccl
|
||||
self.use_custom_allreduce = use_custom_allreduce
|
||||
|
||||
# lazy import to avoid documentation build error
|
||||
# from vllm.distributed.device_communicators.custom_all_reduce import (
|
||||
# CustomAllreduce)
|
||||
from fastvideo.v1.distributed.device_communicators.pynccl import (
|
||||
PyNcclCommunicator)
|
||||
|
||||
self.pynccl_comm: Optional[PyNcclCommunicator] = None
|
||||
if self.world_size > 1:
|
||||
if use_pynccl and self.world_size > 1:
|
||||
self.pynccl_comm = PyNcclCommunicator(
|
||||
group=self.cpu_group,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
# TODO(will): bring in the custom allreduce from vLLM
|
||||
self.ca_comm: Optional[CustomAllreduce] = None
|
||||
if use_custom_allreduce and self.world_size > 1:
|
||||
# Initialize a custom fast all-reduce implementation.
|
||||
self.ca_comm = CustomAllreduce(
|
||||
group=self.cpu_group,
|
||||
device=self.device,
|
||||
)
|
||||
|
||||
def all_reduce(self, input_):
|
||||
# always try custom allreduce first,
|
||||
# and then pynccl.
|
||||
ca_comm = self.ca_comm
|
||||
if ca_comm is not None and not ca_comm.disabled and \
|
||||
ca_comm.should_custom_ar(input_):
|
||||
out = ca_comm.custom_all_reduce(input_)
|
||||
assert out is not None
|
||||
return out
|
||||
pynccl_comm = self.pynccl_comm
|
||||
assert pynccl_comm is not None
|
||||
out = pynccl_comm.all_reduce(input_)
|
||||
@@ -71,6 +102,8 @@ class CudaCommunicator(DeviceCommunicatorBase):
|
||||
torch.distributed.recv(tensor, self.ranks[src], self.device_group)
|
||||
return tensor
|
||||
|
||||
def destroy(self) -> None:
|
||||
def destroy(self):
|
||||
if self.pynccl_comm is not None:
|
||||
self.pynccl_comm = None
|
||||
if self.ca_comm is not None:
|
||||
self.ca_comm = None
|
||||
|
||||
@@ -146,11 +146,11 @@ class PyNcclCommunicator:
|
||||
f"but the input tensor is on {input_tensor.device}")
|
||||
if stream is None:
|
||||
stream = current_stream()
|
||||
self.nccl.ncclAllGather(buffer_type(input_tensor.data_ptr()),
|
||||
buffer_type(output_tensor.data_ptr()),
|
||||
input_tensor.numel(),
|
||||
ncclDataTypeEnum.from_torch(input_tensor.dtype),
|
||||
self.comm, cudaStream_t(stream.cuda_stream))
|
||||
self.nccl.ncclAllGather(
|
||||
buffer_type(input_tensor.data_ptr()),
|
||||
buffer_type(output_tensor.data_ptr()), input_tensor.numel(),
|
||||
ncclDataTypeEnum.from_torch(input_tensor.dtype), self.comm,
|
||||
cudaStream_t(stream.cuda_stream))
|
||||
|
||||
def reduce_scatter(self,
|
||||
output_tensor: torch.Tensor,
|
||||
|
||||
@@ -19,11 +19,9 @@
|
||||
# recompilation of the code every time we want to switch between different
|
||||
# versions. This current implementation, with a **pure** Python wrapper, is
|
||||
# more flexible. We can easily switch between different versions of NCCL by
|
||||
# changing the environment variable `FASTVIDEO_NCCL_SO_PATH`, or the `so_file`
|
||||
# changing the environment variable `VLLM_NCCL_SO_PATH`, or the `so_file`
|
||||
# variable in the code.
|
||||
|
||||
#TODO(will): support FASTVIDEO_NCCL_SO_PATH
|
||||
|
||||
import ctypes
|
||||
import platform
|
||||
from dataclasses import dataclass
|
||||
@@ -35,6 +33,7 @@ from torch.distributed import ReduceOp
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import find_nccl_library
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# === export types and functions from nccl to Python ===
|
||||
@@ -142,7 +141,8 @@ class NCCLLibrary:
|
||||
# note that ncclComm_t is a pointer type, so the first argument
|
||||
# is a pointer to a pointer
|
||||
Function("ncclCommInitRank", ncclResult_t, [
|
||||
ctypes.POINTER(ncclComm_t), ctypes.c_int, ncclUniqueId, ctypes.c_int
|
||||
ctypes.POINTER(ncclComm_t), ctypes.c_int, ncclUniqueId,
|
||||
ctypes.c_int
|
||||
]),
|
||||
# ncclResult_t ncclAllReduce(
|
||||
# const void* sendbuff, void* recvbuff, size_t count,
|
||||
@@ -234,7 +234,7 @@ class NCCLLibrary:
|
||||
"Otherwise, the nccl library might not exist, be corrupted "
|
||||
"or it does not support the current platform %s."
|
||||
"If you already have the library, please set the "
|
||||
"environment variable FASTVIDEO_NCCL_SO_PATH"
|
||||
"environment variable VLLM_NCCL_SO_PATH"
|
||||
" to point to the correct nccl library path.", so_file,
|
||||
platform.platform())
|
||||
raise e
|
||||
@@ -250,7 +250,7 @@ class NCCLLibrary:
|
||||
self._funcs = NCCLLibrary.path_to_dict_mapping[so_file]
|
||||
|
||||
def ncclGetErrorString(self, result: ncclResult_t) -> str:
|
||||
return str(self._funcs["ncclGetErrorString"](result).decode("utf-8"))
|
||||
return self._funcs["ncclGetErrorString"](result).decode("utf-8")
|
||||
|
||||
def NCCL_CHECK(self, result: ncclResult_t) -> None:
|
||||
if result != 0:
|
||||
@@ -269,7 +269,8 @@ class NCCLLibrary:
|
||||
|
||||
def ncclGetUniqueId(self) -> ncclUniqueId:
|
||||
unique_id = ncclUniqueId()
|
||||
self.NCCL_CHECK(self._funcs["ncclGetUniqueId"](ctypes.byref(unique_id)))
|
||||
self.NCCL_CHECK(self._funcs["ncclGetUniqueId"](
|
||||
ctypes.byref(unique_id)))
|
||||
return unique_id
|
||||
|
||||
def ncclCommInitRank(self, world_size: int, unique_id: ncclUniqueId,
|
||||
@@ -316,8 +317,8 @@ class NCCLLibrary:
|
||||
|
||||
def ncclSend(self, sendbuff: buffer_type, count: int, datatype: int,
|
||||
dest: int, comm: ncclComm_t, stream: cudaStream_t) -> None:
|
||||
self.NCCL_CHECK(self._funcs["ncclSend"](sendbuff, count, datatype, dest,
|
||||
comm, stream))
|
||||
self.NCCL_CHECK(self._funcs["ncclSend"](sendbuff, count, datatype,
|
||||
dest, comm, stream))
|
||||
|
||||
def ncclRecv(self, recvbuff: buffer_type, count: int, datatype: int,
|
||||
src: int, comm: ncclComm_t, stream: cudaStream_t) -> None:
|
||||
|
||||
@@ -6,6 +6,7 @@
|
||||
# https://github.com/NVIDIA/Megatron-LM/blob/main/megatron/core/parallel_state.py
|
||||
# Copyright (c) 2022, NVIDIA CORPORATION. All rights reserved.
|
||||
# Adapted from
|
||||
|
||||
"""FastVideo distributed state.
|
||||
It takes over the control of the distributed environment from PyTorch.
|
||||
The typical workflow is:
|
||||
@@ -27,10 +28,11 @@ import gc
|
||||
import pickle
|
||||
import weakref
|
||||
from collections import namedtuple
|
||||
from contextlib import contextmanager
|
||||
from contextlib import contextmanager, nullcontext
|
||||
from dataclasses import dataclass
|
||||
from multiprocessing import shared_memory
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, Union
|
||||
from typing import (TYPE_CHECKING, Any, Callable, Dict, List, Optional, Tuple,
|
||||
Union)
|
||||
from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
@@ -40,12 +42,13 @@ from torch.distributed import Backend, ProcessGroup
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.distributed.device_communicators.base_device_communicator import (
|
||||
DeviceCommunicatorBase)
|
||||
from fastvideo.v1.distributed.device_communicators.cuda_communicator import (
|
||||
CudaCommunicator)
|
||||
from fastvideo.v1.distributed.utils import StatelessProcessGroup
|
||||
from fastvideo.v1.logger import init_logger
|
||||
# from fastvideo.v1.utils import (direct_register_custom_op, resolve_obj_by_qualname,
|
||||
# supports_custom_op)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
from fastvideo.v1.distributed.device_communicators.cuda_communicator import (
|
||||
CudaCommunicator)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -116,6 +119,15 @@ def all_reduce_fake(tensor: torch.Tensor, group_name: str) -> torch.Tensor:
|
||||
return torch.empty_like(tensor)
|
||||
|
||||
|
||||
# if supports_custom_op():
|
||||
# direct_register_custom_op(
|
||||
# op_name="all_reduce",
|
||||
# op_func=all_reduce,
|
||||
# mutates_args=[],
|
||||
# fake_impl=all_reduce_fake,
|
||||
# )
|
||||
|
||||
|
||||
class GroupCoordinator:
|
||||
"""
|
||||
PyTorch ProcessGroup wrapper for a group of processes.
|
||||
@@ -200,11 +212,14 @@ class GroupCoordinator:
|
||||
unique_name=self.unique_name,
|
||||
)
|
||||
|
||||
# from vllm.distributed.device_communicators.shm_broadcast import (
|
||||
# MessageQueue)
|
||||
self.mq_broadcaster = None
|
||||
# if use_message_queue_broadcaster and self.world_size > 1:
|
||||
# self.mq_broadcaster = MessageQueue.create_from_process_group(
|
||||
# self.cpu_group, 1 << 22, 6)
|
||||
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
|
||||
# TODO(will): check if this is needed
|
||||
# self.use_custom_op_call = current_platform.is_cuda_alike()
|
||||
self.use_custom_op_call = False
|
||||
|
||||
@@ -251,13 +266,24 @@ class GroupCoordinator:
|
||||
else:
|
||||
stream = graph_capture_context.stream
|
||||
|
||||
# only cuda uses this function,
|
||||
# so we don't abstract it into the base class
|
||||
maybe_ca_context = nullcontext()
|
||||
from fastvideo.v1.distributed.device_communicators.cuda_communicator import (
|
||||
CudaCommunicator)
|
||||
if self.device_communicator is not None:
|
||||
assert isinstance(self.device_communicator, CudaCommunicator)
|
||||
ca_comm = self.device_communicator.ca_comm
|
||||
if ca_comm is not None:
|
||||
maybe_ca_context = ca_comm.capture() # type: ignore
|
||||
|
||||
# ensure all initialization operations complete before attempting to
|
||||
# capture the graph on another stream
|
||||
curr_stream = torch.cuda.current_stream()
|
||||
if curr_stream != stream:
|
||||
stream.wait_stream(curr_stream)
|
||||
|
||||
with torch.cuda.stream(stream):
|
||||
with torch.cuda.stream(stream), maybe_ca_context:
|
||||
yield graph_capture_context
|
||||
|
||||
def all_reduce(self, input_: torch.Tensor) -> torch.Tensor:
|
||||
@@ -312,15 +338,11 @@ class GroupCoordinator:
|
||||
if world_size == 1:
|
||||
return input_
|
||||
return self.device_communicator.gather(input_, dst, dim)
|
||||
|
||||
def all_to_all_4D(self,
|
||||
input_: torch.Tensor,
|
||||
scatter_dim: int = 2,
|
||||
gather_dim: int = 1) -> torch.Tensor:
|
||||
|
||||
def all_to_all_4D(self, input_: torch.Tensor, scatter_dim: int = 2, gather_dim: int = 1) -> torch.Tensor:
|
||||
if self.world_size == 1:
|
||||
return input_
|
||||
return self.device_communicator.all_to_all_4D(input_, scatter_dim,
|
||||
gather_dim)
|
||||
return self.device_communicator.all_to_all_4D(input_, scatter_dim, gather_dim)
|
||||
|
||||
def broadcast(self, input_: torch.Tensor, src: int = 0):
|
||||
"""Broadcast the input tensor.
|
||||
@@ -416,7 +438,8 @@ class GroupCoordinator:
|
||||
assert src < self.world_size, f"Invalid src rank ({src})"
|
||||
|
||||
assert src != self.rank_in_group, (
|
||||
"Invalid source rank. Source rank is the same as the current rank.")
|
||||
"Invalid source rank. Source rank is the same as the current rank."
|
||||
)
|
||||
|
||||
size_tensor = torch.empty(1, dtype=torch.long, device="cpu")
|
||||
|
||||
@@ -578,7 +601,9 @@ class GroupCoordinator:
|
||||
group=metadata_group)
|
||||
else:
|
||||
# use group for GPU tensors
|
||||
torch.distributed.send(tensor, dst=self.ranks[dst], group=group)
|
||||
torch.distributed.send(tensor,
|
||||
dst=self.ranks[dst],
|
||||
group=group)
|
||||
return None
|
||||
|
||||
def recv_tensor_dict(
|
||||
@@ -669,7 +694,7 @@ class GroupCoordinator:
|
||||
"""NOTE: `src` is the local rank of the source rank."""
|
||||
return self.device_communicator.recv(size, dtype, src)
|
||||
|
||||
def destroy(self) -> None:
|
||||
def destroy(self):
|
||||
if self.device_group is not None:
|
||||
torch.distributed.destroy_process_group(self.device_group)
|
||||
self.device_group = None
|
||||
@@ -730,6 +755,33 @@ def get_tp_group() -> GroupCoordinator:
|
||||
# kept for backward compatibility
|
||||
get_tensor_model_parallel_group = get_tp_group
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@contextmanager
|
||||
def graph_capture(device: torch.device):
|
||||
"""
|
||||
`graph_capture` is a context manager which should surround the code that
|
||||
is capturing the CUDA graph. Its main purpose is to ensure that the
|
||||
some operations will be run after the graph is captured, before the graph
|
||||
is replayed. It returns a `GraphCaptureContext` object which contains the
|
||||
necessary data for the graph capture. Currently, it only contains the
|
||||
stream that the graph capture is running on. This stream is set to the
|
||||
current CUDA stream when the context manager is entered and reset to the
|
||||
default stream when the context manager is exited. This is to ensure that
|
||||
the graph capture is running on a separate stream from the default stream,
|
||||
in order to explicitly distinguish the kernels to capture
|
||||
from other kernels possibly launched on background in the default stream.
|
||||
"""
|
||||
context = GraphCaptureContext(torch.cuda.Stream(device=device))
|
||||
with get_tp_group().graph_capture(context):
|
||||
yield context
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
_ENABLE_CUSTOM_ALL_REDUCE = True
|
||||
|
||||
|
||||
@@ -806,6 +858,7 @@ def initialize_model_parallel(
|
||||
backend = backend or torch.distributed.get_backend(
|
||||
get_world_group().device_group)
|
||||
|
||||
|
||||
num_tensor_model_parallel_groups: int = (world_size //
|
||||
tensor_model_parallel_size)
|
||||
global _TP
|
||||
@@ -826,23 +879,22 @@ def initialize_model_parallel(
|
||||
|
||||
# Build the sequence model-parallel groups.
|
||||
num_sequence_model_parallel_groups: int = (world_size //
|
||||
sequence_model_parallel_size)
|
||||
sequence_model_parallel_size)
|
||||
global _SP
|
||||
assert _SP is None, ("sequence model parallel group is already initialized")
|
||||
group_ranks = []
|
||||
|
||||
|
||||
# Since SP is incompatible with TP and PP, we can use a simpler group creation logic
|
||||
for i in range(num_sequence_model_parallel_groups):
|
||||
# Create groups of consecutive ranks
|
||||
ranks = list(
|
||||
range(i * sequence_model_parallel_size,
|
||||
(i + 1) * sequence_model_parallel_size))
|
||||
ranks = list(range(i * sequence_model_parallel_size,
|
||||
(i + 1) * sequence_model_parallel_size))
|
||||
group_ranks.append(ranks)
|
||||
|
||||
_SP = init_model_parallel_group(group_ranks,
|
||||
get_world_group().local_rank,
|
||||
backend,
|
||||
group_name="sp")
|
||||
get_world_group().local_rank,
|
||||
backend,
|
||||
group_name="sp")
|
||||
|
||||
|
||||
def get_sequence_model_parallel_world_size():
|
||||
@@ -868,7 +920,8 @@ def ensure_model_parallel_initialized(
|
||||
get_world_group().device_group)
|
||||
if not model_parallel_is_initialized():
|
||||
initialize_model_parallel(tensor_model_parallel_size,
|
||||
sequence_model_parallel_size, backend)
|
||||
sequence_model_parallel_size,
|
||||
backend)
|
||||
return
|
||||
|
||||
assert (
|
||||
@@ -876,7 +929,7 @@ def ensure_model_parallel_initialized(
|
||||
), ("tensor parallel group already initialized, but of unexpected size: "
|
||||
f"{get_tensor_model_parallel_world_size()=} vs. "
|
||||
f"{tensor_model_parallel_size=}")
|
||||
|
||||
|
||||
if sequence_model_parallel_size > 1:
|
||||
sp_world_size = get_sp_group().world_size
|
||||
assert (sp_world_size == sequence_model_parallel_size), (
|
||||
@@ -885,9 +938,12 @@ def ensure_model_parallel_initialized(
|
||||
f"{sequence_model_parallel_size=}")
|
||||
|
||||
|
||||
def model_parallel_is_initialized() -> bool:
|
||||
|
||||
def model_parallel_is_initialized():
|
||||
"""Check if tensor, sequence parallel groups are initialized."""
|
||||
return _TP is not None and _SP is not None
|
||||
if _TP is None or _SP is None:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
_TP_STATE_PATCHED = False
|
||||
@@ -918,17 +974,17 @@ def patch_tensor_parallel_group(tp_group: GroupCoordinator):
|
||||
_TP = old_tp_group
|
||||
|
||||
|
||||
def get_tensor_model_parallel_world_size() -> int:
|
||||
def get_tensor_model_parallel_world_size():
|
||||
"""Return world size for the tensor model parallel group."""
|
||||
return get_tp_group().world_size
|
||||
|
||||
|
||||
def get_tensor_model_parallel_rank() -> int:
|
||||
def get_tensor_model_parallel_rank():
|
||||
"""Return my rank for the tensor model parallel group."""
|
||||
return get_tp_group().rank_in_group
|
||||
|
||||
|
||||
def destroy_model_parallel() -> None:
|
||||
def destroy_model_parallel():
|
||||
"""Set the groups to none and destroy them."""
|
||||
global _TP
|
||||
if _TP:
|
||||
@@ -941,7 +997,8 @@ def destroy_model_parallel() -> None:
|
||||
_SP = None
|
||||
|
||||
|
||||
def destroy_distributed_environment() -> None:
|
||||
|
||||
def destroy_distributed_environment():
|
||||
global _WORLD
|
||||
if _WORLD:
|
||||
_WORLD.destroy()
|
||||
@@ -1055,9 +1112,10 @@ def in_the_same_node_as(pg: Union[ProcessGroup, StatelessProcessGroup],
|
||||
|
||||
|
||||
def initialize_tensor_parallel_group(
|
||||
tensor_model_parallel_size: int = 1,
|
||||
backend: Optional[str] = None,
|
||||
group_name_suffix: str = "") -> GroupCoordinator:
|
||||
tensor_model_parallel_size: int = 1,
|
||||
backend: Optional[str] = None,
|
||||
group_name_suffix: str = ""
|
||||
) -> GroupCoordinator:
|
||||
"""Initialize a tensor parallel group for a specific model.
|
||||
|
||||
This function creates a tensor parallel group that can be used with the
|
||||
@@ -1098,8 +1156,7 @@ def initialize_tensor_parallel_group(
|
||||
f"World size ({world_size}) must be divisible by tensor_model_parallel_size ({tensor_model_parallel_size})"
|
||||
|
||||
# Build the tensor model-parallel groups.
|
||||
num_tensor_model_parallel_groups: int = (world_size //
|
||||
tensor_model_parallel_size)
|
||||
num_tensor_model_parallel_groups: int = (world_size // tensor_model_parallel_size)
|
||||
tp_group_ranks = []
|
||||
for i in range(num_tensor_model_parallel_groups):
|
||||
ranks = list(
|
||||
@@ -1110,18 +1167,19 @@ def initialize_tensor_parallel_group(
|
||||
# Create TP group coordinator with a unique name
|
||||
group_name = f"tp_{group_name_suffix}" if group_name_suffix else "tp"
|
||||
tp_group = init_model_parallel_group(tp_group_ranks,
|
||||
get_world_group().local_rank,
|
||||
backend,
|
||||
use_message_queue_broadcaster=True,
|
||||
group_name=group_name)
|
||||
get_world_group().local_rank,
|
||||
backend,
|
||||
use_message_queue_broadcaster=True,
|
||||
group_name=group_name)
|
||||
|
||||
return tp_group
|
||||
|
||||
|
||||
def initialize_sequence_parallel_group(
|
||||
sequence_model_parallel_size: int = 1,
|
||||
backend: Optional[str] = None,
|
||||
group_name_suffix: str = "") -> GroupCoordinator:
|
||||
sequence_model_parallel_size: int = 1,
|
||||
backend: Optional[str] = None,
|
||||
group_name_suffix: str = ""
|
||||
) -> GroupCoordinator:
|
||||
"""Initialize a sequence parallel group for a specific model.
|
||||
|
||||
This function creates a sequence parallel group that can be used with the
|
||||
@@ -1162,22 +1220,132 @@ def initialize_sequence_parallel_group(
|
||||
f"World size ({world_size}) must be divisible by sequence_model_parallel_size ({sequence_model_parallel_size})"
|
||||
|
||||
# Build the sequence model-parallel groups.
|
||||
num_sequence_model_parallel_groups: int = (world_size //
|
||||
sequence_model_parallel_size)
|
||||
num_sequence_model_parallel_groups: int = (world_size // sequence_model_parallel_size)
|
||||
sp_group_ranks = []
|
||||
|
||||
|
||||
for i in range(num_sequence_model_parallel_groups):
|
||||
# Create groups of consecutive ranks
|
||||
ranks = list(
|
||||
range(i * sequence_model_parallel_size,
|
||||
(i + 1) * sequence_model_parallel_size))
|
||||
ranks = list(range(i * sequence_model_parallel_size,
|
||||
(i + 1) * sequence_model_parallel_size))
|
||||
sp_group_ranks.append(ranks)
|
||||
|
||||
# Create SP group coordinator with a unique name
|
||||
group_name = f"sp_{group_name_suffix}" if group_name_suffix else "sp"
|
||||
sp_group = init_model_parallel_group(sp_group_ranks,
|
||||
get_world_group().local_rank,
|
||||
backend,
|
||||
group_name=group_name)
|
||||
get_world_group().local_rank,
|
||||
backend,
|
||||
group_name=group_name)
|
||||
|
||||
return sp_group
|
||||
|
||||
|
||||
_SP_STATE_PATCHED = False
|
||||
|
||||
@contextmanager
|
||||
def patch_sequence_parallel_group(sp_group: GroupCoordinator):
|
||||
"""Patch the sp group temporarily until this function ends.
|
||||
|
||||
This method allows running a model with SP while another model uses TP.
|
||||
|
||||
Args:
|
||||
sp_group (GroupCoordinator): the sp group coordinator
|
||||
"""
|
||||
global _SP_STATE_PATCHED
|
||||
assert not _SP_STATE_PATCHED, "Should not call when it's already patched"
|
||||
|
||||
_SP_STATE_PATCHED = True
|
||||
old_sp_group = get_sp_group()
|
||||
global _SP
|
||||
_SP = sp_group
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
# restore the original state
|
||||
_SP_STATE_PATCHED = False
|
||||
_SP = old_sp_group
|
||||
|
||||
|
||||
|
||||
|
||||
# Example of how to use the independent parallelism functions
|
||||
"""
|
||||
Here's a complete example of how to use the independent parallelism functions
|
||||
for different models:
|
||||
|
||||
```python
|
||||
import torch
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
init_distributed_environment,
|
||||
initialize_tensor_parallel_group,
|
||||
initialize_sequence_parallel_group,
|
||||
patch_tensor_parallel_group,
|
||||
patch_sequence_parallel_group,
|
||||
destroy_model_parallel,
|
||||
destroy_distributed_environment
|
||||
)
|
||||
|
||||
# Initialize the distributed environment
|
||||
init_distributed_environment(
|
||||
world_size=8,
|
||||
rank=torch.distributed.get_rank(),
|
||||
distributed_init_method="tcp://localhost:12345",
|
||||
local_rank=torch.distributed.get_rank() % torch.cuda.device_count()
|
||||
)
|
||||
|
||||
try:
|
||||
# Create a tensor parallel group for model1 with TP=4
|
||||
tp_group_model1 = initialize_tensor_parallel_group(
|
||||
tensor_model_parallel_size=4,
|
||||
group_name_suffix="model1"
|
||||
)
|
||||
|
||||
# Create a sequence parallel group for model2 with SP=2
|
||||
sp_group_model2 = initialize_sequence_parallel_group(
|
||||
sequence_model_parallel_size=2,
|
||||
group_name_suffix="model2"
|
||||
)
|
||||
|
||||
# Create another tensor parallel group for model3 with TP=2
|
||||
tp_group_model3 = initialize_tensor_parallel_group(
|
||||
tensor_model_parallel_size=2,
|
||||
group_name_suffix="model3"
|
||||
)
|
||||
|
||||
# Use model1 with tensor parallelism
|
||||
with patch_tensor_parallel_group(tp_group_model1):
|
||||
# Inside this context, get_tp_group() returns tp_group_model1
|
||||
# Run model1 with tensor parallelism
|
||||
output1 = model1(input1)
|
||||
|
||||
# Use model2 with sequence parallelism
|
||||
with patch_sequence_parallel_group(sp_group_model2):
|
||||
# Inside this context, get_sp_group() returns sp_group_model2
|
||||
# Run model2 with sequence parallelism
|
||||
output2 = model2(input2)
|
||||
|
||||
# Use model3 with a different tensor parallelism configuration
|
||||
with patch_tensor_parallel_group(tp_group_model3):
|
||||
# Inside this context, get_tp_group() returns tp_group_model3
|
||||
# Run model3 with tensor parallelism
|
||||
output3 = model3(input3)
|
||||
|
||||
# You can switch between models as needed
|
||||
with patch_tensor_parallel_group(tp_group_model1):
|
||||
# Back to using model1
|
||||
more_output1 = model1(more_input1)
|
||||
|
||||
finally:
|
||||
# Clean up
|
||||
destroy_model_parallel()
|
||||
destroy_distributed_environment()
|
||||
```
|
||||
|
||||
This approach allows you to:
|
||||
1. Create separate parallel groups for each model
|
||||
2. Use different parallelism strategies for different models
|
||||
3. Switch between models as needed
|
||||
4. Use unique group names to avoid conflicts
|
||||
|
||||
Note that each model can use its own optimal parallelism strategy without
|
||||
interfering with other models.
|
||||
"""
|
||||
|
||||
@@ -12,20 +12,25 @@ from collections import deque
|
||||
from typing import Any, Deque, Dict, Optional, Sequence, Tuple
|
||||
|
||||
import torch
|
||||
from torch.distributed import TCPStore
|
||||
from torch.distributed import ProcessGroup, TCPStore
|
||||
from torch.distributed.distributed_c10d import (Backend, PrefixStore,
|
||||
_get_default_timeout,
|
||||
is_nccl_available)
|
||||
from torch.distributed.rendezvous import rendezvous
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def ensure_divisibility(numerator, denominator) -> None:
|
||||
def ensure_divisibility(numerator, denominator):
|
||||
"""Ensure that numerator is divisible by the denominator."""
|
||||
assert numerator % denominator == 0, "{} is not divisible by {}".format(
|
||||
numerator, denominator)
|
||||
|
||||
|
||||
def divide(numerator: int, denominator: int) -> int:
|
||||
def divide(numerator, denominator):
|
||||
"""Ensure that numerator is divisible by the denominator and return
|
||||
the division value."""
|
||||
ensure_divisibility(numerator, denominator)
|
||||
@@ -57,7 +62,10 @@ def split_tensor_along_last_dim(
|
||||
if contiguous_split_chunks:
|
||||
return tuple(chunk.contiguous() for chunk in tensor_list)
|
||||
|
||||
return tuple(tensor_list)
|
||||
return tensor_list
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
@@ -80,13 +88,17 @@ class StatelessProcessGroup:
|
||||
default_factory=dict)
|
||||
|
||||
# A deque to store the data entries, with key and timestamp.
|
||||
entries: Deque[Tuple[str, float]] = dataclasses.field(default_factory=deque)
|
||||
entries: Deque[Tuple[str,
|
||||
float]] = dataclasses.field(default_factory=deque)
|
||||
|
||||
def __post_init__(self):
|
||||
assert self.rank < self.world_size
|
||||
self.send_dst_counter = {i: 0 for i in range(self.world_size)}
|
||||
self.recv_src_counter = {i: 0 for i in range(self.world_size)}
|
||||
self.broadcast_recv_src_counter = {i: 0 for i in range(self.world_size)}
|
||||
self.broadcast_recv_src_counter = {
|
||||
i: 0
|
||||
for i in range(self.world_size)
|
||||
}
|
||||
|
||||
def send_obj(self, obj: Any, dst: int):
|
||||
"""Send an object to a destination rank."""
|
||||
@@ -96,7 +108,7 @@ class StatelessProcessGroup:
|
||||
self.send_dst_counter[dst] += 1
|
||||
self.entries.append((key, time.time()))
|
||||
|
||||
def expire_data(self) -> None:
|
||||
def expire_data(self):
|
||||
"""Expire data that is older than `data_expiration_seconds` seconds."""
|
||||
while self.entries:
|
||||
# check the oldest entry
|
||||
@@ -110,7 +122,8 @@ class StatelessProcessGroup:
|
||||
def recv_obj(self, src: int) -> Any:
|
||||
"""Receive an object from a source rank."""
|
||||
obj = pickle.loads(
|
||||
self.store.get(f"send_to/{self.rank}/{self.recv_src_counter[src]}"))
|
||||
self.store.get(
|
||||
f"send_to/{self.rank}/{self.recv_src_counter[src]}"))
|
||||
self.recv_src_counter[src] += 1
|
||||
return obj
|
||||
|
||||
|
||||
@@ -1,27 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/types.py
|
||||
|
||||
import argparse
|
||||
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
class CLISubcommand:
|
||||
"""Base class for CLI subcommands"""
|
||||
|
||||
def __init__(self):
|
||||
self.name = ""
|
||||
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
"""Execute the command with the given arguments"""
|
||||
raise NotImplementedError
|
||||
|
||||
def validate(self, args: argparse.Namespace) -> None:
|
||||
"""Validate the arguments for this command"""
|
||||
pass
|
||||
|
||||
def subparser_init(
|
||||
self,
|
||||
subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
|
||||
"""Initialize the subparser for this command"""
|
||||
raise NotImplementedError
|
||||
@@ -1,91 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/serve.py
|
||||
|
||||
import argparse
|
||||
from typing import List
|
||||
|
||||
from fastvideo.v1.entrypoints.cli import utils
|
||||
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
class GenerateSubcommand(CLISubcommand):
|
||||
"""The `generate` subcommand for the FastVideo CLI"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.name = "generate"
|
||||
super().__init__()
|
||||
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
excluded_args = [
|
||||
'subparser', 'config', 'num_gpus', 'master_port',
|
||||
'dispatch_function'
|
||||
]
|
||||
|
||||
# Create a filtered dictionary of arguments
|
||||
filtered_args = {
|
||||
k: v
|
||||
for k, v in vars(args).items()
|
||||
if k not in excluded_args and v is not None
|
||||
}
|
||||
|
||||
main_args = []
|
||||
|
||||
for key, value in filtered_args.items():
|
||||
# Convert underscores to dashes in argument names
|
||||
arg_name = f"--{key.replace('_', '-')}"
|
||||
|
||||
# Handle boolean flags
|
||||
if isinstance(value, bool):
|
||||
if value:
|
||||
main_args.append(arg_name)
|
||||
else:
|
||||
main_args.append(arg_name)
|
||||
main_args.append(str(value))
|
||||
|
||||
utils.launch_distributed(args.num_gpus,
|
||||
main_args,
|
||||
master_port=args.master_port)
|
||||
|
||||
def validate(self, args: argparse.Namespace) -> None:
|
||||
if args.num_gpus is not None and args.num_gpus <= 0:
|
||||
raise ValueError("Number of gpus must be positive")
|
||||
|
||||
if args.master_port is not None and (args.master_port < 1024
|
||||
or args.master_port > 65535):
|
||||
raise ValueError("Master port must be between 1024 and 65535")
|
||||
|
||||
def subparser_init(
|
||||
self,
|
||||
subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
|
||||
generate_parser = subparsers.add_parser(
|
||||
"generate",
|
||||
help="Run inference on a model",
|
||||
usage=
|
||||
"fastvideo generate --model-path MODEL_PATH_OR_ID --prompt PROMPT [OPTIONS]"
|
||||
)
|
||||
|
||||
generate_parser.add_argument(
|
||||
"--config",
|
||||
type=str,
|
||||
default='',
|
||||
required=False,
|
||||
help="Read CLI options from a config YAML file.")
|
||||
|
||||
generate_parser.add_argument("--num-gpus",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of GPUs to use")
|
||||
generate_parser.add_argument("--master-port",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Port for the master process")
|
||||
|
||||
generate_parser = InferenceArgs.add_cli_args(generate_parser)
|
||||
|
||||
return generate_parser
|
||||
|
||||
|
||||
def cmd_init() -> List[CLISubcommand]:
|
||||
return [GenerateSubcommand()]
|
||||
@@ -1,39 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/main.py
|
||||
|
||||
from typing import List
|
||||
|
||||
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.v1.entrypoints.cli.generate import cmd_init as generate_cmd_init
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
def cmd_init() -> List[CLISubcommand]:
|
||||
"""Initialize all commands from separate modules"""
|
||||
commands = []
|
||||
commands.extend(generate_cmd_init())
|
||||
return commands
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = FlexibleArgumentParser(description="FastVideo CLI")
|
||||
parser.add_argument('-v', '--version', action='version', version='0.1.0')
|
||||
|
||||
subparsers = parser.add_subparsers(required=False, dest="subparser")
|
||||
|
||||
cmds = {}
|
||||
for cmd in cmd_init():
|
||||
cmd.subparser_init(subparsers).set_defaults(dispatch_function=cmd.cmd)
|
||||
cmds[cmd.name] = cmd
|
||||
args = parser.parse_args()
|
||||
if args.subparser in cmds:
|
||||
cmds[args.subparser].validate(args)
|
||||
|
||||
if hasattr(args, "dispatch_function"):
|
||||
args.dispatch_function(args)
|
||||
else:
|
||||
parser.print_help()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,57 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def launch_distributed(num_gpus=None, args=None, master_port=None):
|
||||
"""
|
||||
Launch a distributed job with the given arguments
|
||||
|
||||
Args:
|
||||
num_gpus: Number of GPUs to use
|
||||
args: Arguments to pass to v1_fastvideo_inference.py (defaults to sys.argv[1:])
|
||||
master_port: Port for the master process (default: random)
|
||||
"""
|
||||
|
||||
current_env = os.environ.copy()
|
||||
python_executable = sys.executable
|
||||
project_root = os.path.abspath(
|
||||
os.path.join(os.path.dirname(__file__), "../../../.."))
|
||||
main_script = os.path.join(project_root,
|
||||
"fastvideo/v1/sample/v1_fastvideo_inference.py")
|
||||
|
||||
cmd = [
|
||||
python_executable, "-m", "torch.distributed.run",
|
||||
f"--nproc_per_node={num_gpus}"
|
||||
]
|
||||
|
||||
if master_port is not None:
|
||||
cmd.append(f"--master_port={master_port}")
|
||||
|
||||
cmd.append(main_script)
|
||||
cmd.extend(args)
|
||||
|
||||
logger.info("Running inference with %d GPU(s)", num_gpus)
|
||||
logger.info("Launching command: %s", " ".join(cmd))
|
||||
|
||||
current_env["PYTHONIOENCODING"] = "utf-8"
|
||||
process = subprocess.Popen(cmd,
|
||||
env=current_env,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1,
|
||||
encoding='utf-8',
|
||||
errors='replace')
|
||||
|
||||
if process.stdout:
|
||||
for line in iter(process.stdout.readline, ''):
|
||||
print(line.strip())
|
||||
|
||||
return process.wait()
|
||||
+18
-20
@@ -1,8 +1,14 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/envs.py
|
||||
|
||||
# Adapted from vllm
|
||||
# https://github.com/vllm-project/vllm/blob/b382a7f28f739f3b120e5495fd029089d0399428/vllm/envs.py
|
||||
# Copyright 2023 The vLLM Authors.
|
||||
# Copyright 2023 The FastVideo Authors.
|
||||
|
||||
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, Optional
|
||||
import tempfile
|
||||
from typing import TYPE_CHECKING, Any, Callable, Dict, List, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
FASTVIDEO_RINGBUFFER_WARNING_INTERVAL: int = 60
|
||||
@@ -20,7 +26,6 @@ if TYPE_CHECKING:
|
||||
FASTVIDEO_LOGGING_CONFIG_PATH: Optional[str] = None
|
||||
FASTVIDEO_TRACE_FUNCTION: int = 0
|
||||
FASTVIDEO_ATTENTION_BACKEND: Optional[str] = None
|
||||
FASTVIDEO_ATTENTION_CONFIG: Optional[str] = None
|
||||
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "fork"
|
||||
FASTVIDEO_TARGET_DEVICE: str = "cuda"
|
||||
MAX_JOBS: Optional[str] = None
|
||||
@@ -30,14 +35,14 @@ if TYPE_CHECKING:
|
||||
FASTVIDEO_SERVER_DEV_MODE: bool = False
|
||||
|
||||
|
||||
def get_default_cache_root() -> str:
|
||||
def get_default_cache_root():
|
||||
return os.getenv(
|
||||
"XDG_CACHE_HOME",
|
||||
os.path.join(os.path.expanduser("~"), ".cache"),
|
||||
)
|
||||
|
||||
|
||||
def get_default_config_root() -> str:
|
||||
def get_default_config_root():
|
||||
return os.getenv(
|
||||
"XDG_CONFIG_HOME",
|
||||
os.path.join(os.path.expanduser("~"), ".config"),
|
||||
@@ -129,15 +134,13 @@ environment_variables: Dict[str, Callable[[], Any]] = {
|
||||
|
||||
# flag to control if fastvideo should use triton flash attention
|
||||
"FASTVIDEO_USE_TRITON_FLASH_ATTN":
|
||||
lambda:
|
||||
(os.environ.get("FASTVIDEO_USE_TRITON_FLASH_ATTN", "True").lower() in
|
||||
("true", "1")),
|
||||
lambda: (os.environ.get("FASTVIDEO_USE_TRITON_FLASH_ATTN", "True").lower() in
|
||||
("true", "1")),
|
||||
|
||||
# Force fastvideo to use a specific flash-attention version (2 or 3), only valid
|
||||
# when using the flash-attention backend.
|
||||
"FASTVIDEO_FLASH_ATTN_VERSION":
|
||||
lambda: maybe_convert_int(
|
||||
os.environ.get("FASTVIDEO_FLASH_ATTN_VERSION", None)),
|
||||
lambda: maybe_convert_int(os.environ.get("FASTVIDEO_FLASH_ATTN_VERSION", None)),
|
||||
|
||||
# Internal flag to enable Dynamo fullgraph capture
|
||||
"FASTVIDEO_TEST_DYNAMO_FULLGRAPH_CAPTURE":
|
||||
@@ -184,16 +187,12 @@ environment_variables: Dict[str, Callable[[], Any]] = {
|
||||
# Available options:
|
||||
# - "TORCH_SDPA": use torch.nn.MultiheadAttention
|
||||
# - "FLASH_ATTN": use FlashAttention
|
||||
# - "STA" : use sliding tile attention
|
||||
# - "XFORMERS": use XFormers
|
||||
# - "ROCM_FLASH": use ROCmFlashAttention
|
||||
# - "FLASHINFER": use flashinfer
|
||||
"FASTVIDEO_ATTENTION_BACKEND":
|
||||
lambda: os.getenv("FASTVIDEO_ATTENTION_BACKEND", None),
|
||||
|
||||
# Path to the attention configuration file. Only used for sliding tile
|
||||
# attention for now.
|
||||
"FASTVIDEO_ATTENTION_CONFIG":
|
||||
lambda: (None if os.getenv("FASTVIDEO_ATTENTION_CONFIG", None) is None else
|
||||
os.path.expanduser(os.getenv("FASTVIDEO_ATTENTION_CONFIG", "."))),
|
||||
|
||||
# Use dedicated multiprocess context for workers.
|
||||
# Both spawn and fork work
|
||||
"FASTVIDEO_WORKER_MULTIPROC_METHOD":
|
||||
@@ -202,9 +201,8 @@ environment_variables: Dict[str, Callable[[], Any]] = {
|
||||
# Enables torch profiler if set. Path to the directory where torch profiler
|
||||
# traces are saved. Note that it must be an absolute path.
|
||||
"FASTVIDEO_TORCH_PROFILER_DIR":
|
||||
lambda: (None
|
||||
if os.getenv("FASTVIDEO_TORCH_PROFILER_DIR", None) is None else os.
|
||||
path.expanduser(os.getenv("FASTVIDEO_TORCH_PROFILER_DIR", "."))),
|
||||
lambda: (None if os.getenv("FASTVIDEO_TORCH_PROFILER_DIR", None) is None else os
|
||||
.path.expanduser(os.getenv("FASTVIDEO_TORCH_PROFILER_DIR", "."))),
|
||||
|
||||
# If set, fastvideo will run in development mode, which will enable
|
||||
# some additional endpoints for developing and debugging,
|
||||
|
||||
@@ -1,102 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/forward_context.py
|
||||
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.v1.attention import AttentionMetadata
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# TODO(will): check if this is needed
|
||||
# track_batchsize: bool = envs.FASTVIDEO_LOG_BATCHSIZE_INTERVAL >= 0
|
||||
track_batchsize: bool = False
|
||||
last_logging_time: float = 0
|
||||
forward_start_time: float = 0
|
||||
# batchsize_logging_interval: float = envs.FASTVIDEO_LOG_BATCHSIZE_INTERVAL
|
||||
batchsize_logging_interval: float = 1000
|
||||
batchsize_forward_time: defaultdict = defaultdict(list)
|
||||
|
||||
|
||||
#
|
||||
@dataclass
|
||||
class ForwardContext:
|
||||
# TODO(will): check this arg
|
||||
# copy from vllm_config.compilation_config.static_forward_context
|
||||
# attn_layers: Dict[str, Any]
|
||||
# TODO: extend to support per-layer dynamic forward context
|
||||
attn_metadata: "AttentionMetadata" # set dynamically for each forward pass
|
||||
|
||||
|
||||
_forward_context: Optional[ForwardContext] = None
|
||||
|
||||
|
||||
def get_forward_context() -> ForwardContext:
|
||||
"""Get the current forward context."""
|
||||
assert _forward_context is not None, (
|
||||
"Forward context is not set. "
|
||||
"Please use `set_forward_context` to set the forward context.")
|
||||
return _forward_context
|
||||
|
||||
|
||||
# TODO(will): finalize the interface
|
||||
@contextmanager
|
||||
def set_forward_context(current_timestep,
|
||||
attn_metadata,
|
||||
inference_args: InferenceArgs = None):
|
||||
"""A context manager that stores the current forward context,
|
||||
can be attention metadata, etc.
|
||||
Here we can inject common logic for every model forward pass.
|
||||
"""
|
||||
global forward_start_time
|
||||
need_to_track_batchsize = track_batchsize and attn_metadata is not None
|
||||
if need_to_track_batchsize:
|
||||
forward_start_time = time.perf_counter()
|
||||
global _forward_context
|
||||
prev_context = _forward_context
|
||||
_forward_context = ForwardContext(attn_metadata=attn_metadata)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
global last_logging_time, batchsize_logging_interval
|
||||
if need_to_track_batchsize:
|
||||
if hasattr(attn_metadata, "num_prefill_tokens"):
|
||||
# for v0 attention backends
|
||||
batchsize = attn_metadata.num_prefill_tokens + \
|
||||
attn_metadata.num_decode_tokens
|
||||
else:
|
||||
# for v1 attention backends
|
||||
batchsize = attn_metadata.num_input_tokens
|
||||
# we use synchronous scheduling right now,
|
||||
# adding a sync point here should not affect
|
||||
# scheduling of the next batch
|
||||
torch.cuda.synchronize()
|
||||
now = time.perf_counter()
|
||||
# time measurement is in milliseconds
|
||||
batchsize_forward_time[batchsize].append(
|
||||
(now - forward_start_time) * 1000)
|
||||
if now - last_logging_time > batchsize_logging_interval:
|
||||
last_logging_time = now
|
||||
forward_stats = []
|
||||
for bs, times in batchsize_forward_time.items():
|
||||
if len(times) <= 1:
|
||||
# can be cudagraph / profiling run
|
||||
continue
|
||||
medium = torch.quantile(torch.tensor(times), q=0.5).item()
|
||||
medium = round(medium, 2)
|
||||
forward_stats.append((bs, len(times), medium))
|
||||
forward_stats.sort(key=lambda x: x[1], reverse=True)
|
||||
if forward_stats:
|
||||
logger.info(("Batchsize forward time stats "
|
||||
"(batchsize, count, median_time(ms)): %s"),
|
||||
forward_stats)
|
||||
_forward_context = prev_context
|
||||
@@ -1,28 +1,42 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Inspired by SGLang: https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/server_args.py
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
# Adapted from SGLang server_args.py
|
||||
# Copyright 2024-2025 FastVideo Team
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
"""The arguments of FastVideo Inference."""
|
||||
|
||||
import argparse
|
||||
import dataclasses
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
from typing import List, Optional
|
||||
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser
|
||||
|
||||
|
||||
|
||||
@dataclasses.dataclass
|
||||
class InferenceArgs:
|
||||
# Model and path configuration
|
||||
model_path: str
|
||||
|
||||
|
||||
# HuggingFace specific parameters
|
||||
trust_remote_code: bool = False
|
||||
revision: Optional[str] = None
|
||||
|
||||
|
||||
# Parallelism
|
||||
tp_size: int = 1
|
||||
sp_size: int = 1
|
||||
dist_timeout: Optional[int] = None # timeout for torch.distributed
|
||||
|
||||
|
||||
# Video generation parameters
|
||||
height: int = 720
|
||||
width: int = 1280
|
||||
@@ -34,44 +48,48 @@ class InferenceArgs:
|
||||
flow_shift: int = 7
|
||||
|
||||
output_type: str = "pil"
|
||||
|
||||
|
||||
# Model configuration
|
||||
precision: str = "bf16"
|
||||
|
||||
# VAE configuration
|
||||
|
||||
# VAE configurationi
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = False
|
||||
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precision: str = "fp16"
|
||||
text_len: int = 256
|
||||
hidden_state_skip_layer: int = 2
|
||||
|
||||
|
||||
# Secondary text encoder
|
||||
text_encoder_precision_2: str = "fp16"
|
||||
text_len_2: int = 77
|
||||
|
||||
|
||||
# Flow Matching parameters
|
||||
flow_solver: str = "euler"
|
||||
denoise_type: str = "flow"
|
||||
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
|
||||
# Scheduler options
|
||||
scheduler_type: str = "euler"
|
||||
|
||||
|
||||
neg_prompt: Optional[str] = None
|
||||
num_videos: int = 1
|
||||
fps: int = 24
|
||||
use_cpu_offload: bool = False
|
||||
disable_autocast: bool = False
|
||||
|
||||
|
||||
|
||||
# Logging
|
||||
log_level: str = "info"
|
||||
|
||||
|
||||
# Kernel backend
|
||||
attention_backend: Optional[str] = None
|
||||
|
||||
# Inference parameters
|
||||
prompt: Optional[str] = None
|
||||
prompt_path: Optional[str] = None
|
||||
@@ -84,14 +102,28 @@ class InferenceArgs:
|
||||
pass
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
def add_cli_args(parser: argparse.ArgumentParser):
|
||||
parser.add_argument(
|
||||
"--use-v1-text-encoder",
|
||||
action="store_true",
|
||||
help="Use the v1 text encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-v1-vae",
|
||||
action="store_true",
|
||||
help="Use the v1 vae",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-v1-transformer",
|
||||
action="store_true",
|
||||
help="Use the v1 transformer",
|
||||
)
|
||||
# Model and path configuration
|
||||
parser.add_argument(
|
||||
"--model-path",
|
||||
type=str,
|
||||
help="The path of the model weights. This can be a local folder or a Hugging Face repo ID.",
|
||||
required=True,
|
||||
help=
|
||||
"The path of the model weights. This can be a local folder or a Hugging Face repo ID.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dit-weight",
|
||||
@@ -103,7 +135,7 @@ class InferenceArgs:
|
||||
type=str,
|
||||
help="Directory containing StepVideo model",
|
||||
)
|
||||
|
||||
|
||||
# HuggingFace specific parameters
|
||||
parser.add_argument(
|
||||
"--trust-remote-code",
|
||||
@@ -115,8 +147,7 @@ class InferenceArgs:
|
||||
"--revision",
|
||||
type=str,
|
||||
default=InferenceArgs.revision,
|
||||
help=
|
||||
"The specific model version to use (can be a branch name, tag name, or commit id)",
|
||||
help="The specific model version to use (can be a branch name, tag name, or commit id)",
|
||||
)
|
||||
|
||||
# Parallelism
|
||||
@@ -198,6 +229,7 @@ class InferenceArgs:
|
||||
choices=["pil"],
|
||||
help="Output type for the generated video",
|
||||
)
|
||||
|
||||
|
||||
parser.add_argument(
|
||||
"--precision",
|
||||
@@ -206,7 +238,7 @@ class InferenceArgs:
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for the model",
|
||||
)
|
||||
|
||||
|
||||
# VAE configuration
|
||||
parser.add_argument(
|
||||
"--vae-precision",
|
||||
@@ -226,6 +258,7 @@ class InferenceArgs:
|
||||
action="store_true",
|
||||
help="Enable VAE spatial parallelism",
|
||||
)
|
||||
|
||||
|
||||
parser.add_argument(
|
||||
"--text-encoder-precision",
|
||||
@@ -255,7 +288,7 @@ class InferenceArgs:
|
||||
default=InferenceArgs.text_len_2,
|
||||
help="Maximum secondary text length",
|
||||
)
|
||||
|
||||
|
||||
# Flow Matching parameters
|
||||
parser.add_argument(
|
||||
"--flow-solver",
|
||||
@@ -269,7 +302,7 @@ class InferenceArgs:
|
||||
default=InferenceArgs.denoise_type,
|
||||
help="Denoise type for noised inputs",
|
||||
)
|
||||
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
parser.add_argument(
|
||||
"--mask-strategy-file-path",
|
||||
@@ -279,10 +312,9 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--enable-torch-compile",
|
||||
action="store_true",
|
||||
help=
|
||||
"Use torch.compile for speeding up STA inference without teacache",
|
||||
help="Use torch.compile for speeding up STA inference without teacache",
|
||||
)
|
||||
|
||||
|
||||
# Scheduler options
|
||||
parser.add_argument(
|
||||
"--scheduler-type",
|
||||
@@ -290,7 +322,7 @@ class InferenceArgs:
|
||||
default=InferenceArgs.scheduler_type,
|
||||
help="Type of scheduler to use",
|
||||
)
|
||||
|
||||
|
||||
# HunYuan specific parameters
|
||||
parser.add_argument(
|
||||
"--neg-prompt",
|
||||
@@ -318,9 +350,11 @@ class InferenceArgs:
|
||||
parser.add_argument(
|
||||
"--disable-autocast",
|
||||
action="store_true",
|
||||
help=
|
||||
"Disable autocast for denoising loop and vae decoding in pipeline sampling",
|
||||
help="Disable autocast for denoising loop and vae decoding in pipeline sampling",
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
# Logging
|
||||
parser.add_argument(
|
||||
@@ -330,19 +364,26 @@ class InferenceArgs:
|
||||
help="The logging level of all loggers.",
|
||||
)
|
||||
|
||||
# Kernel backend
|
||||
parser.add_argument(
|
||||
"--attention-backend",
|
||||
type=str,
|
||||
choices=["flashinfer", "triton", "torch_native"],
|
||||
default=InferenceArgs.attention_backend,
|
||||
help="Choose the kernels for attention layers.",
|
||||
)
|
||||
|
||||
# Inference parameters
|
||||
prompt_group = parser.add_mutually_exclusive_group(required=True)
|
||||
prompt_group.add_argument(
|
||||
parser.add_argument(
|
||||
"--prompt",
|
||||
type=str,
|
||||
help="Text prompt for video generation",
|
||||
)
|
||||
prompt_group.add_argument(
|
||||
parser.add_argument(
|
||||
"--prompt-path",
|
||||
type=str,
|
||||
help="Path to a text file containing the prompt",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--output-path",
|
||||
type=str,
|
||||
@@ -356,20 +397,21 @@ class InferenceArgs:
|
||||
help="Random seed for reproducibility",
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "InferenceArgs":
|
||||
def from_cli_args(cls, args: argparse.Namespace):
|
||||
args.tp_size = args.tensor_parallel_size
|
||||
args.sp_size = args.sequence_parallel_size
|
||||
args.flow_shift = getattr(args, "shift", args.flow_shift)
|
||||
|
||||
|
||||
# Get all fields from the dataclass
|
||||
attrs = [attr.name for attr in dataclasses.fields(cls)]
|
||||
|
||||
|
||||
# Create a dictionary of attribute values, with defaults for missing attributes
|
||||
kwargs = {}
|
||||
for attr in attrs:
|
||||
# Convert snake_case attribute name to kebab-case CLI argument name
|
||||
cli_attr = attr.replace('_', '-')
|
||||
|
||||
# Handle renamed attributes or those with multiple CLI names
|
||||
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
|
||||
kwargs[attr] = args.tensor_parallel_size
|
||||
@@ -381,24 +423,20 @@ class InferenceArgs:
|
||||
else:
|
||||
default_value = getattr(cls, attr, None)
|
||||
kwargs[attr] = getattr(args, attr, default_value)
|
||||
|
||||
|
||||
return cls(**kwargs)
|
||||
|
||||
def check_inference_args(self) -> None:
|
||||
"""Validate inference arguments for consistency"""
|
||||
|
||||
def check_inference_args(self):
|
||||
"""Validate inference arguments for consistency"""
|
||||
|
||||
# Validate VAE spatial parallelism with VAE tiling
|
||||
if self.vae_sp and not self.vae_tiling:
|
||||
raise ValueError(
|
||||
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
|
||||
)
|
||||
if self.prompt_path and not self.prompt_path.endswith(".txt"):
|
||||
raise ValueError("prompt_path must be a text file")
|
||||
|
||||
raise ValueError("Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True.")
|
||||
assert self.prompt is not None or self.prompt_path is not None, "Either prompt or prompt_path must be provided"
|
||||
assert self.prompt_path.endswith(".txt"), "prompt_path must be a text file"
|
||||
|
||||
_inference_args = None
|
||||
|
||||
|
||||
def prepare_inference_args(argv: List[str]) -> InferenceArgs:
|
||||
"""
|
||||
Prepare the inference arguments from the command line arguments.
|
||||
@@ -419,18 +457,15 @@ def prepare_inference_args(argv: List[str]) -> InferenceArgs:
|
||||
_inference_args = inference_args
|
||||
return inference_args
|
||||
|
||||
|
||||
def get_inference_args() -> InferenceArgs:
|
||||
global _inference_args
|
||||
if _inference_args is None:
|
||||
raise ValueError("Inference arguments not set")
|
||||
return _inference_args
|
||||
|
||||
|
||||
class DeprecatedAction(argparse.Action):
|
||||
|
||||
def __init__(self, option_strings, dest, nargs=0, **kwargs):
|
||||
super().__init__(option_strings, dest, nargs=nargs, **kwargs)
|
||||
super(DeprecatedAction, self).__init__(
|
||||
option_strings, dest, nargs=nargs, **kwargs
|
||||
)
|
||||
|
||||
def __call__(self, parser, namespace, values, option_string=None):
|
||||
raise ValueError(self.help)
|
||||
raise ValueError(self.help)
|
||||
@@ -1,30 +1,31 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Inference module for diffusion models.
|
||||
|
||||
This module provides classes and functions for running inference with diffusion models.
|
||||
"""
|
||||
|
||||
import os
|
||||
import time
|
||||
import torch
|
||||
from typing import Any, Dict
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.pipelines import ComposedPipelineBase, build_pipeline
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines import (ComposedPipelineBase, ForwardBatch,
|
||||
build_pipeline)
|
||||
# TODO(will): remove, check if this is hunyuan specific
|
||||
from fastvideo.v1.utils import align_to
|
||||
# TODO(will): remove, move this to hunyuan stage
|
||||
from fastvideo.v1.pipelines.implementations.hunyuan.constants import NEGATIVE_PROMPT
|
||||
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class InferenceEngine:
|
||||
"""
|
||||
Engine for running inference with diffusion models.
|
||||
"""
|
||||
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
pipeline: ComposedPipelineBase,
|
||||
@@ -40,7 +41,9 @@ class InferenceEngine:
|
||||
"""
|
||||
self.pipeline = pipeline
|
||||
self.inference_args = inference_args
|
||||
|
||||
# TODO(will): this is a hack to get the default negative prompt
|
||||
self.default_negative_prompt = NEGATIVE_PROMPT
|
||||
|
||||
@classmethod
|
||||
def create_engine(
|
||||
cls,
|
||||
@@ -64,7 +67,7 @@ class InferenceEngine:
|
||||
is not recognized.
|
||||
"""
|
||||
|
||||
logger.info("Building pipeline...")
|
||||
logger.info(f"Building pipeline...")
|
||||
|
||||
# TODO(will): I don't really like this api.
|
||||
# it should be something closer to pipeline_cls.from_pretrained(...)
|
||||
@@ -72,11 +75,12 @@ class InferenceEngine:
|
||||
# checkpoint_path) and have it handle everything.
|
||||
# TODO(Peiyuan): Then maybe we should only pass in model path and device, not the entire inference args?
|
||||
pipeline = build_pipeline(inference_args)
|
||||
logger.info("Pipeline Ready")
|
||||
logger.info(f"Pipeline Ready")
|
||||
|
||||
|
||||
# Create the inference engine
|
||||
return cls(pipeline, inference_args)
|
||||
|
||||
|
||||
def run(
|
||||
self,
|
||||
prompt: str,
|
||||
@@ -94,7 +98,7 @@ class InferenceEngine:
|
||||
Returns:
|
||||
A dictionary containing the generated videos and metadata.
|
||||
"""
|
||||
out_dict: Dict[str, Any] = dict()
|
||||
out_dict = dict()
|
||||
|
||||
num_videos_per_prompt = inference_args.num_videos
|
||||
seed = inference_args.seed
|
||||
@@ -107,6 +111,8 @@ class InferenceEngine:
|
||||
flow_shift = inference_args.flow_shift
|
||||
embedded_guidance_scale = inference_args.embedded_cfg_scale
|
||||
|
||||
|
||||
|
||||
# ========================================================================
|
||||
# Arguments: target_width, target_height, target_video_length
|
||||
# ========================================================================
|
||||
@@ -115,8 +121,9 @@ class InferenceEngine:
|
||||
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})")
|
||||
|
||||
target_height = align_to(height, 16)
|
||||
target_width = align_to(width, 16)
|
||||
@@ -128,13 +135,16 @@ class InferenceEngine:
|
||||
# Arguments: prompt, new_prompt, negative_prompt
|
||||
# ========================================================================
|
||||
if not isinstance(prompt, str):
|
||||
raise TypeError(
|
||||
f"`prompt` must be a string, but got {type(prompt)}")
|
||||
raise TypeError(f"`prompt` must be a string, but got {type(prompt)}")
|
||||
prompt = prompt.strip()
|
||||
|
||||
# negative prompt
|
||||
if negative_prompt is not None:
|
||||
negative_prompt = negative_prompt.strip()
|
||||
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)}")
|
||||
negative_prompt = negative_prompt.strip()
|
||||
|
||||
|
||||
# TODO(PY): move to hunyuan stage
|
||||
latents_size = [(video_length - 1) // 4 + 1, height // 8, width // 8]
|
||||
@@ -191,13 +201,13 @@ class InferenceEngine:
|
||||
samples = self.pipeline.forward(
|
||||
batch=batch,
|
||||
inference_args=inference_args,
|
||||
).output
|
||||
)[0]
|
||||
# TODO(will): fix and move to hunyuan stage
|
||||
# out_dict["seeds"] = batch.seeds
|
||||
out_dict["samples"] = samples
|
||||
out_dict["prompts"] = prompt
|
||||
|
||||
gen_time = time.time() - start_time
|
||||
logger.info("Success, time: %s", gen_time)
|
||||
logger.info(f"Success, time: {gen_time}")
|
||||
|
||||
return out_dict
|
||||
return out_dict
|
||||
@@ -1,17 +1,55 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/activation.py
|
||||
"""Custom activation functions."""
|
||||
import math
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
# TODO (will): remove this dependency
|
||||
from fastvideo.v1.layers.custom_op import CustomOp
|
||||
from vllm.distributed import (divide, get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size)
|
||||
from vllm.model_executor.custom_op import CustomOp
|
||||
from vllm.model_executor.utils import set_weight_attrs
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
|
||||
|
||||
@CustomOp.register("fatrelu_and_mul")
|
||||
class FatreluAndMul(CustomOp):
|
||||
"""An activation function for FATReLU.
|
||||
|
||||
The function computes x -> FATReLU(x[:d]) * x[d:] where
|
||||
d = x.shape[-1] // 2.
|
||||
This is used in openbmb/MiniCPM-S-1B-sft.
|
||||
|
||||
Shapes:
|
||||
x: (num_tokens, 2 * d) or (batch_size, seq_len, 2 * d)
|
||||
return: (num_tokens, d) or (batch_size, seq_len, d)
|
||||
"""
|
||||
|
||||
def __init__(self, threshold: float = 0.):
|
||||
super().__init__()
|
||||
self.threshold = threshold
|
||||
if current_platform.is_cuda_alike():
|
||||
self.op = torch.ops._C.fatrelu_and_mul
|
||||
elif current_platform.is_cpu():
|
||||
self._forward_method = self.forward_native
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
d = x.shape[-1] // 2
|
||||
x1 = x[..., :d]
|
||||
x2 = x[..., d:]
|
||||
x1 = F.threshold(x1, self.threshold, 0.0)
|
||||
return x1 * x2
|
||||
|
||||
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
|
||||
d = x.shape[-1] // 2
|
||||
output_shape = (x.shape[:-1] + (d, ))
|
||||
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||
self.op(out, x, self.threshold)
|
||||
return out
|
||||
|
||||
|
||||
@CustomOp.register("silu_and_mul")
|
||||
class SiluAndMul(CustomOp):
|
||||
"""An activation function for SwiGLU.
|
||||
@@ -27,6 +65,9 @@ class SiluAndMul(CustomOp):
|
||||
super().__init__()
|
||||
if current_platform.is_cuda_alike() or current_platform.is_cpu():
|
||||
self.op = torch.ops._C.silu_and_mul
|
||||
elif current_platform.is_xpu():
|
||||
from vllm._ipex_ops import ipex_ops
|
||||
self.op = ipex_ops.silu_and_mul
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
@@ -40,6 +81,50 @@ class SiluAndMul(CustomOp):
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
|
||||
d = x.shape[-1] // 2
|
||||
output_shape = (x.shape[:-1] + (d, ))
|
||||
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
|
||||
@CustomOp.register("mul_and_silu")
|
||||
class MulAndSilu(CustomOp):
|
||||
"""An activation function for SwiGLU.
|
||||
|
||||
The function computes x -> x[:d] * silu(x[d:]) where d = x.shape[-1] // 2.
|
||||
|
||||
Shapes:
|
||||
x: (num_tokens, 2 * d) or (batch_size, seq_len, 2 * d)
|
||||
return: (num_tokens, d) or (batch_size, seq_len, d)
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
if current_platform.is_cuda_alike():
|
||||
self.op = torch.ops._C.mul_and_silu
|
||||
elif current_platform.is_xpu():
|
||||
from vllm._ipex_ops import ipex_ops
|
||||
self.op = ipex_ops.silu_and_mul
|
||||
elif current_platform.is_cpu():
|
||||
self._forward_method = self.forward_native
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
d = x.shape[-1] // 2
|
||||
return x[..., :d] * F.silu(x[..., d:])
|
||||
|
||||
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
|
||||
d = x.shape[-1] // 2
|
||||
output_shape = (x.shape[:-1] + (d, ))
|
||||
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
# TODO implement forward_xpu for MulAndSilu
|
||||
# def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
|
||||
|
||||
|
||||
@CustomOp.register("gelu_and_mul")
|
||||
class GeluAndMul(CustomOp):
|
||||
@@ -62,6 +147,12 @@ class GeluAndMul(CustomOp):
|
||||
self.op = torch.ops._C.gelu_and_mul
|
||||
elif approximate == "tanh":
|
||||
self.op = torch.ops._C.gelu_tanh_and_mul
|
||||
elif current_platform.is_xpu():
|
||||
from vllm._ipex_ops import ipex_ops
|
||||
if approximate == "none":
|
||||
self.op = ipex_ops.gelu_and_mul
|
||||
else:
|
||||
self.op = ipex_ops.gelu_tanh_and_mul
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
@@ -75,6 +166,13 @@ class GeluAndMul(CustomOp):
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
|
||||
d = x.shape[-1] // 2
|
||||
output_shape = (x.shape[:-1] + (d, ))
|
||||
out = torch.empty(output_shape, dtype=x.dtype, device=x.device)
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
def extra_repr(self) -> str:
|
||||
return f'approximate={repr(self.approximate)}'
|
||||
|
||||
@@ -86,6 +184,9 @@ class NewGELU(CustomOp):
|
||||
super().__init__()
|
||||
if current_platform.is_cuda_alike() or current_platform.is_cpu():
|
||||
self.op = torch.ops._C.gelu_new
|
||||
elif current_platform.is_xpu():
|
||||
from vllm._ipex_ops import ipex_ops
|
||||
self.op = ipex_ops.gelu_new
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
@@ -102,6 +203,31 @@ class NewGELU(CustomOp):
|
||||
return self.op(x)
|
||||
|
||||
|
||||
@CustomOp.register("gelu_fast")
|
||||
class FastGELU(CustomOp):
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
if current_platform.is_cuda_alike() or current_platform.is_cpu():
|
||||
self.op = torch.ops._C.gelu_fast
|
||||
elif current_platform.is_xpu():
|
||||
from vllm._ipex_ops import ipex_ops
|
||||
self.op = ipex_ops.gelu_fast
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
return 0.5 * x * (1.0 + torch.tanh(x * 0.7978845608 *
|
||||
(1.0 + 0.044715 * x * x)))
|
||||
|
||||
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
|
||||
out = torch.empty_like(x)
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.op(x)
|
||||
|
||||
|
||||
@CustomOp.register("quick_gelu")
|
||||
class QuickGELU(CustomOp):
|
||||
# https://github.com/huggingface/transformers/blob/main/src/transformers/activations.py#L90
|
||||
@@ -109,6 +235,9 @@ class QuickGELU(CustomOp):
|
||||
super().__init__()
|
||||
if current_platform.is_cuda_alike() or current_platform.is_cpu():
|
||||
self.op = torch.ops._C.gelu_quick
|
||||
elif current_platform.is_xpu():
|
||||
from vllm._ipex_ops import ipex_ops
|
||||
self.op = ipex_ops.gelu_quick
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
@@ -119,12 +248,78 @@ class QuickGELU(CustomOp):
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
|
||||
out = torch.empty_like(x)
|
||||
self.op(out, x)
|
||||
return out
|
||||
|
||||
# TODO implement forward_xpu for QuickGELU
|
||||
# def forward_xpu(self, x: torch.Tensor) -> torch.Tensor:
|
||||
|
||||
|
||||
@CustomOp.register("relu2")
|
||||
class ReLUSquaredActivation(CustomOp):
|
||||
"""
|
||||
Applies the relu^2 activation introduced in https://arxiv.org/abs/2109.08668v2
|
||||
"""
|
||||
|
||||
def forward_native(self, x: torch.Tensor) -> torch.Tensor:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
return torch.square(F.relu(x))
|
||||
|
||||
def forward_cuda(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.forward_native(x)
|
||||
|
||||
|
||||
class ScaledActivation(nn.Module):
|
||||
"""An activation function with post-scale parameters.
|
||||
|
||||
This is used for some quantization methods like AWQ.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
act_module: nn.Module,
|
||||
intermediate_size: int,
|
||||
input_is_parallel: bool = True,
|
||||
params_dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
super().__init__()
|
||||
self.act = act_module
|
||||
self.input_is_parallel = input_is_parallel
|
||||
if input_is_parallel:
|
||||
tp_size = get_tensor_model_parallel_world_size()
|
||||
intermediate_size_per_partition = divide(intermediate_size,
|
||||
tp_size)
|
||||
else:
|
||||
intermediate_size_per_partition = intermediate_size
|
||||
if params_dtype is None:
|
||||
params_dtype = torch.get_default_dtype()
|
||||
self.scales = nn.Parameter(
|
||||
torch.empty(intermediate_size_per_partition, dtype=params_dtype))
|
||||
set_weight_attrs(self.scales, {"weight_loader": self.weight_loader})
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return self.act(x) / self.scales
|
||||
|
||||
def weight_loader(self, param: nn.Parameter, loaded_weight: torch.Tensor):
|
||||
param_data = param.data
|
||||
if self.input_is_parallel:
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
shard_size = param_data.shape[0]
|
||||
start_idx = tp_rank * shard_size
|
||||
loaded_weight = loaded_weight.narrow(0, start_idx, shard_size)
|
||||
assert param_data.shape == loaded_weight.shape
|
||||
param_data.copy_(loaded_weight)
|
||||
|
||||
|
||||
_ACTIVATION_REGISTRY = {
|
||||
"gelu": nn.GELU,
|
||||
"gelu_fast": FastGELU,
|
||||
"gelu_new": NewGELU,
|
||||
"gelu_pytorch_tanh": lambda: nn.GELU(approximate="tanh"),
|
||||
"relu": nn.ReLU,
|
||||
"relu2": ReLUSquaredActivation,
|
||||
"silu": nn.SiLU,
|
||||
"quick_gelu": QuickGELU,
|
||||
}
|
||||
|
||||
@@ -1,92 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/custom_op.py
|
||||
|
||||
from typing import Any, Callable, Dict, Type
|
||||
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class CustomOp(nn.Module):
|
||||
"""
|
||||
Base class for custom ops.
|
||||
Dispatches the forward method to the appropriate backend.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
self._forward_method = self.dispatch_forward()
|
||||
|
||||
def forward(self, *args, **kwargs) -> Any:
|
||||
return self._forward_method(*args, **kwargs)
|
||||
|
||||
def forward_native(self, *args, **kwargs) -> Any:
|
||||
"""PyTorch-native implementation of the forward method.
|
||||
This method is optional. If implemented, it can be used with compilers
|
||||
such as torch.compile or PyTorch XLA. Also, it can be used for testing
|
||||
purposes.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def forward_cuda(self, *args, **kwargs) -> Any:
|
||||
raise NotImplementedError
|
||||
|
||||
def forward_cpu(self, *args, **kwargs) -> Any:
|
||||
# By default, we assume that CPU ops are compatible with CUDA ops.
|
||||
return self.forward_cuda(*args, **kwargs)
|
||||
|
||||
def forward_tpu(self, *args, **kwargs) -> Any:
|
||||
# By default, we assume that TPU ops are compatible with the
|
||||
# PyTorch-native implementation.
|
||||
# NOTE(woosuk): This is a placeholder for future extensions.
|
||||
return self.forward_native(*args, **kwargs)
|
||||
|
||||
def forward_oot(self, *args, **kwargs) -> Any:
|
||||
# By default, we assume that OOT ops are compatible with the
|
||||
# PyTorch-native implementation.
|
||||
return self.forward_native(*args, **kwargs)
|
||||
|
||||
def dispatch_forward(self) -> Callable:
|
||||
# NOTE(woosuk): Here we assume that vLLM was built for only one
|
||||
# specific backend. Currently, we do not support dynamic dispatching.
|
||||
enabled = self.enabled()
|
||||
|
||||
if not enabled:
|
||||
return self.forward_native
|
||||
|
||||
return self.forward_cuda
|
||||
|
||||
@classmethod
|
||||
def enabled(cls) -> bool:
|
||||
# since we are not using Inductor, we always return True
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def default_on() -> bool:
|
||||
"""
|
||||
On by default if level < CompilationLevel.PIECEWISE
|
||||
Specifying 'all' or 'none' in custom_op takes precedence.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
# Dictionary of all custom ops (classes, indexed by registered name).
|
||||
# To check if an op with a name is enabled, call .enabled() on the class.
|
||||
# Examples:
|
||||
# - MyOp.enabled()
|
||||
# - op_registry["my_op"].enabled()
|
||||
op_registry: Dict[str, Type['CustomOp']] = {}
|
||||
|
||||
# Decorator to register custom ops.
|
||||
@classmethod
|
||||
def register(cls, name: str) -> Callable:
|
||||
|
||||
def decorator(op_cls):
|
||||
assert name not in cls.op_registry, f"Duplicate op name: {name}"
|
||||
op_cls.name = name
|
||||
cls.op_registry[name] = op_cls
|
||||
return op_cls
|
||||
|
||||
return decorator
|
||||
@@ -1,12 +1,11 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/layernorm.py
|
||||
"""Custom normalization layers."""
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.v1.layers.custom_op import CustomOp
|
||||
from vllm.model_executor.custom_op import CustomOp
|
||||
|
||||
|
||||
@CustomOp.register("rms_norm")
|
||||
@@ -21,7 +20,6 @@ class RMSNorm(CustomOp):
|
||||
self,
|
||||
hidden_size: int,
|
||||
eps: float = 1e-6,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
var_hidden_size: Optional[int] = None,
|
||||
has_weight: bool = True,
|
||||
) -> None:
|
||||
@@ -102,22 +100,129 @@ class RMSNorm(CustomOp):
|
||||
)
|
||||
return out
|
||||
|
||||
def forward_hpu(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
residual: Optional[torch.Tensor] = None,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
from vllm_hpu_extension.ops import HPUFusedRMSNorm
|
||||
if HPUFusedRMSNorm is None:
|
||||
return self.forward_native(x, residual)
|
||||
if residual is not None:
|
||||
orig_shape = x.shape
|
||||
residual += x.view(residual.shape)
|
||||
# Note: HPUFusedRMSNorm requires 3D tensors as inputs
|
||||
x = HPUFusedRMSNorm.apply(residual, self.weight,
|
||||
self.variance_epsilon)
|
||||
return x.view(orig_shape), residual
|
||||
|
||||
x = HPUFusedRMSNorm.apply(x, self.weight, self.variance_epsilon)
|
||||
return x
|
||||
|
||||
def forward_xpu(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
residual: Optional[torch.Tensor] = None,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
if self.variance_size_override is not None:
|
||||
return self.forward_native(x, residual)
|
||||
|
||||
from vllm._ipex_ops import ipex_ops as ops
|
||||
|
||||
if residual is not None:
|
||||
ops.fused_add_rms_norm(
|
||||
x,
|
||||
residual,
|
||||
self.weight.data,
|
||||
self.variance_epsilon,
|
||||
)
|
||||
return x, residual
|
||||
return ops.rms_norm(
|
||||
x,
|
||||
self.weight.data,
|
||||
self.variance_epsilon,
|
||||
)
|
||||
|
||||
def extra_repr(self) -> str:
|
||||
s = f"hidden_size={self.weight.data.size(0)}"
|
||||
s += f", eps={self.variance_epsilon}"
|
||||
return s
|
||||
|
||||
|
||||
@CustomOp.register("gemma_rms_norm")
|
||||
class GemmaRMSNorm(CustomOp):
|
||||
"""RMS normalization for Gemma.
|
||||
|
||||
Two differences from the above RMSNorm:
|
||||
1. x * (1 + w) instead of x * w.
|
||||
2. (x * w).to(orig_dtype) instead of x.to(orig_dtype) * w.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
eps: float = 1e-6,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.zeros(hidden_size))
|
||||
self.variance_epsilon = eps
|
||||
|
||||
@staticmethod
|
||||
def forward_static(
|
||||
weight: torch.Tensor,
|
||||
variance_epsilon: float,
|
||||
x: torch.Tensor,
|
||||
residual: Optional[torch.Tensor],
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
orig_dtype = x.dtype
|
||||
if residual is not None:
|
||||
x = x + residual
|
||||
residual = x
|
||||
|
||||
x = x.float()
|
||||
variance = x.pow(2).mean(dim=-1, keepdim=True)
|
||||
x = x * torch.rsqrt(variance + variance_epsilon)
|
||||
# Llama does x.to(float16) * w whilst Gemma is (x * w).to(float16)
|
||||
# See https://github.com/huggingface/transformers/pull/29402
|
||||
x = x * (1.0 + weight.float())
|
||||
x = x.to(orig_dtype)
|
||||
return x if residual is None else (x, residual)
|
||||
|
||||
def forward_native(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
residual: Optional[torch.Tensor] = None,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
"""PyTorch-native implementation equivalent to forward()."""
|
||||
return self.forward_static(self.weight.data, self.variance_epsilon, x,
|
||||
residual)
|
||||
|
||||
def forward_cuda(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
residual: Optional[torch.Tensor] = None,
|
||||
) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]:
|
||||
if torch.compiler.is_compiling():
|
||||
return self.forward_native(x, residual)
|
||||
|
||||
if not getattr(self, "_is_compiled", False):
|
||||
self.forward_static = torch.compile( # type: ignore
|
||||
self.forward_static)
|
||||
self._is_compiled = True
|
||||
return self.forward_native(x, residual)
|
||||
|
||||
|
||||
|
||||
class ScaleResidual(nn.Module):
|
||||
"""
|
||||
Applies gated residual connection.
|
||||
"""
|
||||
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, residual: torch.Tensor, x: torch.Tensor,
|
||||
gate: torch.Tensor) -> torch.Tensor:
|
||||
|
||||
def forward(self, residual: torch.Tensor, x: torch.Tensor, gate: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply gated residual connection."""
|
||||
return residual + x * gate
|
||||
|
||||
@@ -131,7 +236,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
|
||||
This reduces memory bandwidth by combining memory-bound operations.
|
||||
"""
|
||||
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
@@ -142,46 +247,45 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
):
|
||||
super().__init__()
|
||||
if norm_type == "rms":
|
||||
self.norm = RMSNorm(hidden_size,
|
||||
has_weight=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
self.norm = RMSNorm(hidden_size, has_weight=elementwise_affine, eps=eps, dtype=dtype)
|
||||
elif norm_type == "layer":
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
self.norm = nn.LayerNorm(hidden_size, elementwise_affine=elementwise_affine, eps=eps, dtype=dtype)
|
||||
else:
|
||||
raise NotImplementedError(f"Norm type {norm_type} not implemented")
|
||||
|
||||
def forward(self, residual: torch.Tensor, x: torch.Tensor,
|
||||
gate: torch.Tensor, shift: torch.Tensor,
|
||||
scale: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
|
||||
def forward(
|
||||
self,
|
||||
residual: torch.Tensor,
|
||||
x: torch.Tensor,
|
||||
gate: torch.Tensor,
|
||||
shift: torch.Tensor,
|
||||
scale: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Apply gated residual connection, followed by layernorm and
|
||||
scale/shift in a single fused operation.
|
||||
Apply gated residual connection, followed by layernorm and scale/shift in a single fused operation.
|
||||
|
||||
Returns:
|
||||
Tuple containing:
|
||||
- normalized and modulated output
|
||||
- residual value (value after residual connection
|
||||
but before normalization)
|
||||
- residual value (value after residual connection but before normalization)
|
||||
"""
|
||||
# Apply residual connection with gating
|
||||
residual_output = residual + x * gate
|
||||
# Apply normalization
|
||||
normalized = self.norm(residual_output)
|
||||
# Apply scale and shift
|
||||
modulated = normalized * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
modulated = normalized * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
return modulated, residual_output
|
||||
|
||||
|
||||
|
||||
|
||||
class LayerNormScaleShift(nn.Module):
|
||||
"""
|
||||
Fused operation that combines LayerNorm with scale and shift operations.
|
||||
This reduces memory bandwidth by combining memory-bound operations.
|
||||
"""
|
||||
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
@@ -192,19 +296,13 @@ class LayerNormScaleShift(nn.Module):
|
||||
):
|
||||
super().__init__()
|
||||
if norm_type == "rms":
|
||||
self.norm = RMSNorm(hidden_size,
|
||||
has_weight=elementwise_affine,
|
||||
eps=eps)
|
||||
self.norm = RMSNorm(hidden_size, has_weight=elementwise_affine, eps=eps)
|
||||
elif norm_type == "layer":
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
self.norm = nn.LayerNorm(hidden_size, elementwise_affine=elementwise_affine, eps=eps, dtype=dtype)
|
||||
else:
|
||||
raise NotImplementedError(f"Norm type {norm_type} not implemented")
|
||||
|
||||
def forward(self, x: torch.Tensor, shift: torch.Tensor,
|
||||
scale: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply ln followed by scale and shift in a single fused operation."""
|
||||
|
||||
def forward(self, x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply layernorm followed by scale and shift in a single fused operation."""
|
||||
normalized = self.norm(x)
|
||||
return normalized * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
|
||||
|
||||
+280
-53
@@ -1,23 +1,23 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/linear.py
|
||||
|
||||
import itertools
|
||||
from abc import abstractmethod
|
||||
from typing import Optional, Union
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch.nn.parameter import Parameter
|
||||
from torch.nn.parameter import Parameter, UninitializedParameter
|
||||
|
||||
from fastvideo.v1.distributed import (divide, get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size,
|
||||
split_tensor_along_last_dim,
|
||||
tensor_model_parallel_all_gather,
|
||||
tensor_model_parallel_all_reduce)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
# TODO(will): remove this import by copying the definition from vLLM then
|
||||
# manually import each quantization method we want to use. Refer to SGLang
|
||||
from vllm.model_executor.layers.quantization.base_config import (
|
||||
QuantizationConfig, QuantizeMethodBase)
|
||||
|
||||
from fastvideo.v1.distributed import (divide, get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size,
|
||||
split_tensor_along_last_dim,
|
||||
tensor_model_parallel_all_gather,
|
||||
tensor_model_parallel_all_reduce)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
# yapf: disable
|
||||
from fastvideo.v1.models.parameter import (BasevLLMParameter,
|
||||
BlockQuantScaleParameter,
|
||||
@@ -31,18 +31,39 @@ from fastvideo.v1.models.utils import set_weight_attrs
|
||||
logger = init_logger(__name__)
|
||||
|
||||
WEIGHT_LOADER_V2_SUPPORTED = [
|
||||
"CompressedTensorsLinearMethod", "AWQMarlinLinearMethod", "AWQLinearMethod",
|
||||
"GPTQMarlinLinearMethod", "Fp8LinearMethod", "MarlinLinearMethod",
|
||||
"QQQLinearMethod", "GPTQMarlin24LinearMethod", "TPUInt8LinearMethod",
|
||||
"GPTQLinearMethod", "FBGEMMFp8LinearMethod", "ModelOptFp8LinearMethod",
|
||||
"IPEXAWQLinearMethod", "IPEXGPTQLinearMethod", "HQQMarlinMethod",
|
||||
"QuarkLinearMethod"
|
||||
"CompressedTensorsLinearMethod", "AWQMarlinLinearMethod",
|
||||
"AWQLinearMethod", "GPTQMarlinLinearMethod", "Fp8LinearMethod",
|
||||
"MarlinLinearMethod", "QQQLinearMethod", "GPTQMarlin24LinearMethod",
|
||||
"TPUInt8LinearMethod", "GPTQLinearMethod", "FBGEMMFp8LinearMethod",
|
||||
"ModelOptFp8LinearMethod", "IPEXAWQLinearMethod", "IPEXGPTQLinearMethod",
|
||||
"HQQMarlinMethod", "QuarkLinearMethod"
|
||||
]
|
||||
|
||||
|
||||
def adjust_scalar_to_fused_array(
|
||||
param: torch.Tensor, loaded_weight: torch.Tensor,
|
||||
shard_id: Union[str, int]) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
def adjust_marlin_shard(param, shard_size, shard_offset):
|
||||
marlin_tile_size = getattr(param, "marlin_tile_size", None)
|
||||
if marlin_tile_size is None:
|
||||
return shard_size, shard_offset
|
||||
|
||||
return shard_size * marlin_tile_size, shard_offset * marlin_tile_size
|
||||
|
||||
|
||||
def adjust_bitsandbytes_4bit_shard(param: Parameter,
|
||||
shard_offsets: dict[str, tuple[int, int]],
|
||||
loaded_shard_id: str) -> tuple[int, int]:
|
||||
"""Adjust the quantization offsets and sizes for BitsAndBytes sharding."""
|
||||
|
||||
total, _ = shard_offsets["total"]
|
||||
orig_offset, orig_size = shard_offsets[loaded_shard_id]
|
||||
|
||||
quantized_total = param.data.shape[0]
|
||||
quantized_offset = orig_offset * quantized_total // total
|
||||
quantized_size = orig_size * quantized_total // total
|
||||
|
||||
return quantized_size, quantized_offset
|
||||
|
||||
|
||||
def adjust_scalar_to_fused_array(param, loaded_weight, shard_id):
|
||||
"""For fused modules (QKV and MLP) we have an array of length
|
||||
N that holds 1 scale for each "logical" matrix. So the param
|
||||
is an array of length N. The loaded_weight corresponds to
|
||||
@@ -73,7 +94,7 @@ class LinearMethodBase(QuantizeMethodBase):
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int], input_size: int,
|
||||
output_size: int, params_dtype: torch.dtype,
|
||||
**extra_weight_attrs) -> None:
|
||||
**extra_weight_attrs):
|
||||
"""Create weights for a linear layer.
|
||||
The weights will be set as attributes of the layer.
|
||||
|
||||
@@ -106,7 +127,7 @@ class UnquantizedLinearMethod(LinearMethodBase):
|
||||
input_size_per_partition: int,
|
||||
output_partition_sizes: list[int], input_size: int,
|
||||
output_size: int, params_dtype: torch.dtype,
|
||||
**extra_weight_attrs) -> None:
|
||||
**extra_weight_attrs):
|
||||
weight = Parameter(torch.empty(sum(output_partition_sizes),
|
||||
input_size_per_partition,
|
||||
dtype=params_dtype),
|
||||
@@ -213,10 +234,20 @@ class ReplicatedLinear(LinearBase):
|
||||
else:
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
def weight_loader(self, param: Parameter,
|
||||
loaded_weight: torch.Tensor) -> None:
|
||||
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor):
|
||||
# If the weight on disk does not have a shape, give it one
|
||||
# (such scales for AutoFp8).
|
||||
# Special case for GGUF
|
||||
|
||||
is_gguf_weight = getattr(param, "is_gguf_weight", False)
|
||||
is_gguf_weight_type = getattr(param, "is_gguf_weight_type", False)
|
||||
if is_gguf_weight_type:
|
||||
param.weight_type = loaded_weight.item()
|
||||
|
||||
# Materialize GGUF UninitializedParameter
|
||||
if is_gguf_weight and isinstance(param, UninitializedParameter):
|
||||
param.materialize(loaded_weight.shape, dtype=loaded_weight.dtype)
|
||||
|
||||
if len(loaded_weight.shape) == 0:
|
||||
loaded_weight = loaded_weight.reshape(1)
|
||||
|
||||
@@ -307,7 +338,8 @@ class ColumnParallelLinear(LinearBase):
|
||||
in WEIGHT_LOADER_V2_SUPPORTED else self.weight_loader))
|
||||
if bias:
|
||||
self.bias = Parameter(
|
||||
torch.empty(self.output_size_per_partition, dtype=params_dtype))
|
||||
torch.empty(self.output_size_per_partition,
|
||||
dtype=params_dtype))
|
||||
set_weight_attrs(self.bias, {
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
@@ -315,13 +347,30 @@ class ColumnParallelLinear(LinearBase):
|
||||
else:
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
def weight_loader(self, param: Parameter,
|
||||
loaded_weight: torch.Tensor) -> None:
|
||||
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor):
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
output_dim = getattr(param, "output_dim", None)
|
||||
|
||||
is_sharded_weight = getattr(param, "is_sharded_weight", False)
|
||||
is_sharded_weight = is_sharded_weight
|
||||
use_bitsandbytes_4bit = getattr(param, "use_bitsandbytes_4bit", False)
|
||||
# bitsandbytes loads the weights of the specific portion
|
||||
# no need to narrow
|
||||
is_sharded_weight = is_sharded_weight or use_bitsandbytes_4bit
|
||||
|
||||
# Special case for GGUF
|
||||
is_gguf_weight = getattr(param, "is_gguf_weight", False)
|
||||
is_gguf_weight_type = getattr(param, "is_gguf_weight_type", False)
|
||||
if is_gguf_weight_type:
|
||||
param.weight_type = loaded_weight.item()
|
||||
|
||||
# Materialize GGUF UninitializedParameter
|
||||
if is_gguf_weight and isinstance(param, UninitializedParameter):
|
||||
final_shape = list(loaded_weight.shape)
|
||||
if output_dim is not None:
|
||||
tp_size = get_tensor_model_parallel_world_size()
|
||||
assert final_shape[output_dim] % tp_size == 0
|
||||
final_shape[output_dim] = final_shape[output_dim] // tp_size
|
||||
param.materialize(final_shape, dtype=loaded_weight.dtype)
|
||||
|
||||
param_data = param.data
|
||||
if output_dim is not None and not is_sharded_weight:
|
||||
@@ -338,8 +387,7 @@ class ColumnParallelLinear(LinearBase):
|
||||
assert param_data.shape == loaded_weight.shape
|
||||
param_data.copy_(loaded_weight)
|
||||
|
||||
def weight_loader_v2(self, param: Parameter,
|
||||
loaded_weight: torch.Tensor) -> None:
|
||||
def weight_loader_v2(self, param: Parameter, loaded_weight: torch.Tensor):
|
||||
# Special case for loading scales off disk, which often do not
|
||||
# have a shape (such as in the case of AutoFP8).
|
||||
if len(loaded_weight.shape) == 0:
|
||||
@@ -347,9 +395,7 @@ class ColumnParallelLinear(LinearBase):
|
||||
loaded_weight = loaded_weight.reshape(1)
|
||||
param.load_column_parallel_weight(loaded_weight=loaded_weight)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_: torch.Tensor) -> tuple[torch.Tensor, Optional[Parameter]]:
|
||||
def forward(self, input_) -> tuple[torch.Tensor, Optional[Parameter]]:
|
||||
bias = self.bias if not self.skip_bias_add else None
|
||||
|
||||
# Matrix multiply.
|
||||
@@ -419,7 +465,40 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
def weight_loader(self,
|
||||
param: Parameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: Optional[int] = None) -> None:
|
||||
loaded_shard_id: Optional[int] = None):
|
||||
|
||||
# Special case for GGUF
|
||||
# initialize GGUF param after we know the quantize type
|
||||
is_gguf_weight = getattr(param, "is_gguf_weight", False)
|
||||
is_gguf_weight_type = getattr(param, "is_gguf_weight_type", False)
|
||||
if is_gguf_weight_type:
|
||||
if loaded_shard_id is not None:
|
||||
param.data[loaded_shard_id].copy_(loaded_weight)
|
||||
param.shard_weight_type[loaded_shard_id] = loaded_weight.item()
|
||||
else:
|
||||
param.shard_weight_type = {
|
||||
i: loaded_weight.item()
|
||||
for i, _ in enumerate(self.output_sizes)
|
||||
}
|
||||
return
|
||||
|
||||
if is_gguf_weight:
|
||||
tp_size = get_tensor_model_parallel_world_size()
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
|
||||
output_dim = getattr(param, "output_dim", None)
|
||||
shard_size = loaded_weight.size(output_dim) // tp_size
|
||||
start_idx = tp_rank * shard_size
|
||||
|
||||
if loaded_shard_id is not None:
|
||||
loaded_weight = loaded_weight.narrow(output_dim, start_idx,
|
||||
shard_size)
|
||||
param.shard_id.append(loaded_shard_id)
|
||||
param.shard_id_map[loaded_shard_id] = len(param.data_container)
|
||||
param.data_container.append(loaded_weight)
|
||||
if len(param.data_container) == 2:
|
||||
self.qweight = param.materialize_nested()
|
||||
return
|
||||
|
||||
param_data = param.data
|
||||
output_dim = getattr(param, "output_dim", None)
|
||||
@@ -440,11 +519,34 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
param_data.copy_(loaded_weight)
|
||||
return
|
||||
current_shard_offset = 0
|
||||
use_bitsandbytes_4bit = getattr(param, "use_bitsandbytes_4bit",
|
||||
False)
|
||||
shard_offsets: list[tuple[int, int, int]] = []
|
||||
for i, output_size in enumerate(self.output_sizes):
|
||||
shard_offsets.append((i, current_shard_offset, output_size))
|
||||
current_shard_offset += output_size
|
||||
packed_dim = getattr(param, "packed_dim", None)
|
||||
for shard_id, shard_offset, shard_size in shard_offsets:
|
||||
# Special case for Quantization.
|
||||
# If quantized, we need to adjust the offset and size to account
|
||||
# for the packing.
|
||||
if packed_dim == output_dim:
|
||||
shard_size = shard_size // param.pack_factor
|
||||
shard_offset = shard_offset // param.pack_factor
|
||||
# Special case for Marlin.
|
||||
shard_size, shard_offset = adjust_marlin_shard(
|
||||
param, shard_size, shard_offset)
|
||||
|
||||
if use_bitsandbytes_4bit:
|
||||
index = list(itertools.accumulate([0] + self.output_sizes))
|
||||
orig_offsets = {
|
||||
str(i): (index[i], size)
|
||||
for i, size in enumerate(self.output_sizes)
|
||||
}
|
||||
orig_offsets["total"] = (self.output_size, 0)
|
||||
shard_size, shard_offset = adjust_bitsandbytes_4bit_shard(
|
||||
param, orig_offsets, str(shard_id))
|
||||
|
||||
loaded_weight_shard = loaded_weight.narrow(
|
||||
output_dim, shard_offset, shard_size)
|
||||
self.weight_loader(param, loaded_weight_shard, shard_id)
|
||||
@@ -456,13 +558,31 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
if output_dim is not None:
|
||||
shard_offset = sum(self.output_sizes[:loaded_shard_id]) // tp_size
|
||||
shard_size = self.output_sizes[loaded_shard_id] // tp_size
|
||||
# Special case for quantization.
|
||||
# If quantized, we need to adjust the offset and size to account
|
||||
# for the packing.
|
||||
packed_dim = getattr(param, "packed_dim", None)
|
||||
if packed_dim == output_dim:
|
||||
shard_size = shard_size // param.pack_factor
|
||||
shard_offset = shard_offset // param.pack_factor
|
||||
# Special case for Marlin.
|
||||
shard_size, shard_offset = adjust_marlin_shard(
|
||||
param, shard_size, shard_offset)
|
||||
|
||||
use_bitsandbytes_4bit = getattr(param, "use_bitsandbytes_4bit",
|
||||
False)
|
||||
is_sharded_weight = getattr(param, "is_sharded_weight", False)
|
||||
# bitsandbytes loads the weights of the specific portion
|
||||
# no need to narrow
|
||||
is_sharded_weight = is_sharded_weight
|
||||
is_sharded_weight = is_sharded_weight or use_bitsandbytes_4bit
|
||||
|
||||
param_data = param_data.narrow(output_dim, shard_offset, shard_size)
|
||||
if use_bitsandbytes_4bit:
|
||||
shard_size = loaded_weight.shape[output_dim]
|
||||
shard_offset = loaded_weight.shape[output_dim] * \
|
||||
loaded_shard_id
|
||||
|
||||
param_data = param_data.narrow(output_dim, shard_offset,
|
||||
shard_size)
|
||||
start_idx = tp_rank * shard_size
|
||||
if not is_sharded_weight:
|
||||
loaded_weight = loaded_weight.narrow(output_dim, start_idx,
|
||||
@@ -491,7 +611,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
param_data.copy_(loaded_weight)
|
||||
|
||||
def _load_fused_module_from_checkpoint(self, param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor) -> None:
|
||||
loaded_weight: torch.Tensor):
|
||||
"""
|
||||
Handle special case for models where MLP layers are already
|
||||
fused on disk. In this case, we have no shard id. This function
|
||||
@@ -512,22 +632,21 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
|
||||
# Special case for Quantization.
|
||||
# If quantized, we need to adjust the offset and size to account
|
||||
# for the packing.
|
||||
if isinstance(
|
||||
param,
|
||||
(PackedColumnParameter,
|
||||
PackedvLLMParameter)) and param.packed_dim == param.output_dim:
|
||||
if isinstance(param, (PackedColumnParameter, PackedvLLMParameter
|
||||
)) and param.packed_dim == param.output_dim:
|
||||
shard_size, shard_offset = \
|
||||
param.adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size, shard_offset=shard_offset)
|
||||
|
||||
loaded_weight_shard = loaded_weight.narrow(param.output_dim,
|
||||
shard_offset, shard_size)
|
||||
shard_offset,
|
||||
shard_size)
|
||||
self.weight_loader_v2(param, loaded_weight_shard, shard_id)
|
||||
|
||||
def weight_loader_v2(self,
|
||||
param: BasevLLMParameter,
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: Optional[int] = None) -> None:
|
||||
loaded_shard_id: Optional[int] = None):
|
||||
if loaded_shard_id is None:
|
||||
if isinstance(param, PerTensorScaleParameter):
|
||||
param.load_merged_column_weight(loaded_weight=loaded_weight,
|
||||
@@ -615,7 +734,8 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
self.num_heads = divide(self.total_num_heads, tp_size)
|
||||
if tp_size >= self.total_num_kv_heads:
|
||||
self.num_kv_heads = 1
|
||||
self.num_kv_head_replicas = divide(tp_size, self.total_num_kv_heads)
|
||||
self.num_kv_head_replicas = divide(tp_size,
|
||||
self.total_num_kv_heads)
|
||||
else:
|
||||
self.num_kv_heads = divide(self.total_num_kv_heads, tp_size)
|
||||
self.num_kv_head_replicas = 1
|
||||
@@ -637,7 +757,7 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
quant_config=quant_config,
|
||||
prefix=prefix)
|
||||
|
||||
def _get_shard_offset_mapping(self, loaded_shard_id: str) -> Optional[int]:
|
||||
def _get_shard_offset_mapping(self, loaded_shard_id: str):
|
||||
shard_offset_mapping = {
|
||||
"q": 0,
|
||||
"k": self.num_heads * self.head_size,
|
||||
@@ -646,7 +766,7 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
}
|
||||
return shard_offset_mapping.get(loaded_shard_id)
|
||||
|
||||
def _get_shard_size_mapping(self, loaded_shard_id: str) -> Optional[int]:
|
||||
def _get_shard_size_mapping(self, loaded_shard_id: str):
|
||||
shard_size_mapping = {
|
||||
"q": self.num_heads * self.head_size,
|
||||
"k": self.num_kv_heads * self.head_size,
|
||||
@@ -679,16 +799,15 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
# Special case for Quantization.
|
||||
# If quantized, we need to adjust the offset and size to account
|
||||
# for the packing.
|
||||
if isinstance(
|
||||
param,
|
||||
(PackedColumnParameter,
|
||||
PackedvLLMParameter)) and param.packed_dim == param.output_dim:
|
||||
if isinstance(param, (PackedColumnParameter, PackedvLLMParameter
|
||||
)) and param.packed_dim == param.output_dim:
|
||||
shard_size, shard_offset = \
|
||||
param.adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size, shard_offset=shard_offset)
|
||||
|
||||
loaded_weight_shard = loaded_weight.narrow(param.output_dim,
|
||||
shard_offset, shard_size)
|
||||
shard_offset,
|
||||
shard_size)
|
||||
self.weight_loader_v2(param, loaded_weight_shard, shard_id)
|
||||
|
||||
def weight_loader_v2(self,
|
||||
@@ -722,6 +841,40 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
loaded_weight: torch.Tensor,
|
||||
loaded_shard_id: Optional[str] = None):
|
||||
|
||||
# Special case for GGUF
|
||||
# initialize GGUF param after we know the quantize type
|
||||
is_gguf_weight = getattr(param, "is_gguf_weight", False)
|
||||
is_gguf_weight_type = getattr(param, "is_gguf_weight_type", False)
|
||||
if is_gguf_weight_type:
|
||||
idx_map = {"q": 0, "k": 1, "v": 2}
|
||||
if loaded_shard_id is not None:
|
||||
param.data[idx_map[loaded_shard_id]].copy_(loaded_weight)
|
||||
param.shard_weight_type[loaded_shard_id] = loaded_weight.item()
|
||||
else:
|
||||
param.shard_weight_type = {
|
||||
k: loaded_weight.item()
|
||||
for k in idx_map
|
||||
}
|
||||
return
|
||||
|
||||
if is_gguf_weight:
|
||||
tp_size = get_tensor_model_parallel_world_size()
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
|
||||
output_dim = getattr(param, "output_dim", None)
|
||||
shard_size = loaded_weight.size(output_dim) // tp_size
|
||||
start_idx = tp_rank * shard_size
|
||||
|
||||
if loaded_shard_id is not None:
|
||||
loaded_weight = loaded_weight.narrow(output_dim, start_idx,
|
||||
shard_size)
|
||||
param.shard_id.append(loaded_shard_id)
|
||||
param.shard_id_map[loaded_shard_id] = len(param.data_container)
|
||||
param.data_container.append(loaded_weight)
|
||||
if len(param.data_container) == 3:
|
||||
self.qweight = param.materialize_nested()
|
||||
return
|
||||
|
||||
param_data = param.data
|
||||
output_dim = getattr(param, "output_dim", None)
|
||||
# Special case for AQLM codebooks.
|
||||
@@ -749,8 +902,38 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
("v", (self.total_num_heads + self.total_num_kv_heads) *
|
||||
self.head_size, self.total_num_kv_heads * self.head_size),
|
||||
]
|
||||
use_bitsandbytes_4bit = getattr(param, "use_bitsandbytes_4bit",
|
||||
False)
|
||||
|
||||
packed_dim = getattr(param, "packed_dim", None)
|
||||
for shard_id, shard_offset, shard_size in shard_offsets:
|
||||
# Special case for Quantized Weights.
|
||||
# If quantized, we need to adjust the offset and size to account
|
||||
# for the packing.
|
||||
if packed_dim == output_dim:
|
||||
shard_size = shard_size // param.pack_factor
|
||||
shard_offset = shard_offset // param.pack_factor
|
||||
|
||||
# Special case for Marlin.
|
||||
shard_size, shard_offset = adjust_marlin_shard(
|
||||
param, shard_size, shard_offset)
|
||||
|
||||
if use_bitsandbytes_4bit:
|
||||
orig_qkv_offsets = {
|
||||
"q": (0, self.total_num_heads * self.head_size),
|
||||
"k": (self.total_num_heads * self.head_size,
|
||||
self.total_num_kv_heads * self.head_size),
|
||||
"v":
|
||||
((self.total_num_heads + self.total_num_kv_heads) *
|
||||
self.head_size,
|
||||
self.total_num_kv_heads * self.head_size),
|
||||
"total":
|
||||
((self.total_num_heads + 2 * self.total_num_kv_heads) *
|
||||
self.head_size, 0)
|
||||
}
|
||||
|
||||
shard_size, shard_offset = adjust_bitsandbytes_4bit_shard(
|
||||
param, orig_qkv_offsets, shard_id)
|
||||
|
||||
loaded_weight_shard = loaded_weight.narrow(
|
||||
output_dim, shard_offset, shard_size)
|
||||
@@ -772,13 +955,42 @@ class QKVParallelLinear(ColumnParallelLinear):
|
||||
shard_offset = (self.num_heads +
|
||||
self.num_kv_heads) * self.head_size
|
||||
shard_size = self.num_kv_heads * self.head_size
|
||||
# Special case for Quantized Weights.
|
||||
# If quantized, we need to adjust the offset and size to account
|
||||
# for the packing.
|
||||
packed_dim = getattr(param, "packed_dim", None)
|
||||
if packed_dim == output_dim:
|
||||
shard_size = shard_size // param.pack_factor
|
||||
shard_offset = shard_offset // param.pack_factor
|
||||
|
||||
# Special case for Marlin.
|
||||
shard_size, shard_offset = adjust_marlin_shard(
|
||||
param, shard_size, shard_offset)
|
||||
|
||||
use_bitsandbytes_4bit = getattr(param, "use_bitsandbytes_4bit",
|
||||
False)
|
||||
is_sharded_weight = getattr(param, "is_sharded_weight", False)
|
||||
# bitsandbytes loads the weights of the specific portion
|
||||
# no need to narrow
|
||||
is_sharded_weight = is_sharded_weight
|
||||
is_sharded_weight = is_sharded_weight or use_bitsandbytes_4bit
|
||||
|
||||
param_data = param_data.narrow(output_dim, shard_offset, shard_size)
|
||||
if use_bitsandbytes_4bit:
|
||||
orig_qkv_offsets = {
|
||||
"q": (0, self.num_heads * self.head_size),
|
||||
"k": (self.num_heads * self.head_size,
|
||||
self.num_kv_heads * self.head_size),
|
||||
"v":
|
||||
((self.num_heads + self.num_kv_heads) * self.head_size,
|
||||
self.num_kv_heads * self.head_size),
|
||||
"total":
|
||||
((self.num_heads + 2 * self.num_kv_heads) * self.head_size,
|
||||
0)
|
||||
}
|
||||
shard_size, shard_offset = adjust_bitsandbytes_4bit_shard(
|
||||
param, orig_qkv_offsets, loaded_shard_id)
|
||||
|
||||
param_data = param_data.narrow(output_dim, shard_offset,
|
||||
shard_size)
|
||||
if loaded_shard_id == "q":
|
||||
shard_id = tp_rank
|
||||
else:
|
||||
@@ -888,11 +1100,26 @@ class RowParallelLinear(LinearBase):
|
||||
|
||||
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor):
|
||||
tp_rank = get_tensor_model_parallel_rank()
|
||||
tp_size = get_tensor_model_parallel_world_size()
|
||||
input_dim = getattr(param, "input_dim", None)
|
||||
use_bitsandbytes_4bit = getattr(param, "use_bitsandbytes_4bit", False)
|
||||
is_sharded_weight = getattr(param, "is_sharded_weight", False)
|
||||
# bitsandbytes loads the weights of the specific portion
|
||||
# no need to narrow
|
||||
is_sharded_weight = is_sharded_weight
|
||||
is_sharded_weight = is_sharded_weight or use_bitsandbytes_4bit
|
||||
|
||||
# Special case for GGUF
|
||||
is_gguf_weight = getattr(param, "is_gguf_weight", False)
|
||||
is_gguf_weight_type = getattr(param, "is_gguf_weight_type", False)
|
||||
if is_gguf_weight_type:
|
||||
param.weight_type = loaded_weight.item()
|
||||
|
||||
# Materialize GGUF UninitializedParameter
|
||||
if is_gguf_weight and isinstance(param, UninitializedParameter):
|
||||
weight_shape = list(loaded_weight.shape)
|
||||
if input_dim:
|
||||
weight_shape[input_dim] = weight_shape[input_dim] // tp_size
|
||||
param.materialize(tuple(weight_shape), dtype=loaded_weight.dtype)
|
||||
|
||||
param_data = param.data
|
||||
if input_dim is not None and not is_sharded_weight:
|
||||
|
||||
+17
-17
@@ -1,12 +1,8 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
from typing import Optional
|
||||
|
||||
|
||||
class MLP(nn.Module):
|
||||
@@ -14,12 +10,12 @@ class MLP(nn.Module):
|
||||
MLP for DiT blocks, NO gated linear units
|
||||
TODO: add Tensor Parallel
|
||||
"""
|
||||
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
input_dim: int,
|
||||
mlp_hidden_dim: int,
|
||||
output_dim: Optional[int] = None,
|
||||
output_dim: int = None,
|
||||
bias: bool = True,
|
||||
act_type: str = "gelu_pytorch_tanh",
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
@@ -27,20 +23,24 @@ class MLP(nn.Module):
|
||||
super().__init__()
|
||||
self.fc_in = ReplicatedLinear(
|
||||
input_dim,
|
||||
mlp_hidden_dim, # For activation func like SiLU that need 2x width
|
||||
mlp_hidden_dim, # For activation functions like SiLU that need 2x width
|
||||
bias=bias,
|
||||
params_dtype=dtype)
|
||||
|
||||
params_dtype=dtype
|
||||
)
|
||||
|
||||
self.act = get_act_fn(act_type)
|
||||
if output_dim is None:
|
||||
output_dim = input_dim
|
||||
self.fc_out = ReplicatedLinear(mlp_hidden_dim,
|
||||
output_dim,
|
||||
bias=bias,
|
||||
params_dtype=dtype)
|
||||
|
||||
self.fc_out = ReplicatedLinear(
|
||||
mlp_hidden_dim,
|
||||
output_dim,
|
||||
bias=bias,
|
||||
params_dtype=dtype
|
||||
)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x, _ = self.fc_in(x)
|
||||
x = self.act(x)
|
||||
x, _ = self.fc_out(x)
|
||||
return x
|
||||
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/rotary_embedding.py
|
||||
|
||||
# Adapted from
|
||||
# https://github.com/huggingface/transformers/blob/v4.33.2/src/transformers/models/llama/modeling_llama.py
|
||||
@@ -23,13 +22,15 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Rotary Positional Embeddings."""
|
||||
import math
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
from vllm.model_executor.custom_op import CustomOp
|
||||
from fastvideo.v1.distributed.parallel_state import get_sp_group
|
||||
from fastvideo.v1.layers.custom_op import CustomOp
|
||||
|
||||
|
||||
def _rotate_neox(x: torch.Tensor) -> torch.Tensor:
|
||||
x1 = x[..., :x.shape[-1] // 2]
|
||||
@@ -167,32 +168,111 @@ class RotaryEmbedding(CustomOp):
|
||||
# are in-place operations that update the query and key tensors.
|
||||
if offsets is not None:
|
||||
ops.batched_rotary_embedding(positions, query, key, self.head_size,
|
||||
self.cos_sin_cache, self.is_neox_style,
|
||||
self.rotary_dim, offsets)
|
||||
self.cos_sin_cache,
|
||||
self.is_neox_style, self.rotary_dim,
|
||||
offsets)
|
||||
else:
|
||||
ops.rotary_embedding(positions, query, key, self.head_size,
|
||||
self.cos_sin_cache, self.is_neox_style)
|
||||
return query, key
|
||||
|
||||
def forward_xpu(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
offsets: Optional[torch.Tensor] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
from vllm._ipex_ops import ipex_ops as ops
|
||||
|
||||
self.cos_sin_cache = self.cos_sin_cache.to(positions.device,
|
||||
dtype=query.dtype)
|
||||
# ops.rotary_embedding()/batched_rotary_embedding()
|
||||
# are in-place operations that update the query and key tensors.
|
||||
if offsets is not None:
|
||||
ops.batched_rotary_embedding(positions, query, key, self.head_size,
|
||||
self.cos_sin_cache,
|
||||
self.is_neox_style, self.rotary_dim,
|
||||
offsets)
|
||||
else:
|
||||
ops.rotary_embedding(positions, query, key, self.head_size,
|
||||
self.cos_sin_cache, self.is_neox_style)
|
||||
return query, key
|
||||
|
||||
def forward_hpu(
|
||||
self,
|
||||
positions: torch.Tensor,
|
||||
query: torch.Tensor,
|
||||
key: torch.Tensor,
|
||||
offsets: Optional[torch.Tensor] = None,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
from habana_frameworks.torch.hpex.kernels import (
|
||||
RotaryPosEmbeddingMode, apply_rotary_pos_emb)
|
||||
if offsets is not None:
|
||||
offsets = offsets.view(positions.shape[0], -1)
|
||||
positions = positions + offsets
|
||||
positions = positions.flatten()
|
||||
num_tokens = positions.shape[0]
|
||||
cos_sin = self.cos_sin_cache.index_select(0, positions).view(
|
||||
num_tokens, 1, -1)
|
||||
cos, sin = cos_sin.chunk(2, dim=-1)
|
||||
# HPU RoPE kernel requires hidden dimension for cos and sin to be equal
|
||||
# to query hidden dimension, so the original tensors need to be
|
||||
# expanded
|
||||
# GPT-NeoX kernel requires position_ids = None, offset, mode = BLOCKWISE
|
||||
# and expansion of cos/sin tensors via concatenation
|
||||
# GPT-J kernel requires position_ids = None, offset = 0, mode = PAIRWISE
|
||||
# and expansion of cos/sin tensors via repeat_interleave
|
||||
rope_mode: RotaryPosEmbeddingMode
|
||||
if self.is_neox_style:
|
||||
rope_mode = RotaryPosEmbeddingMode.BLOCKWISE
|
||||
cos = torch.cat((cos, cos), dim=-1)
|
||||
sin = torch.cat((sin, sin), dim=-1)
|
||||
else:
|
||||
rope_mode = RotaryPosEmbeddingMode.PAIRWISE
|
||||
sin = torch.repeat_interleave(sin,
|
||||
2,
|
||||
dim=-1,
|
||||
output_size=cos_sin.shape[-1])
|
||||
cos = torch.repeat_interleave(cos,
|
||||
2,
|
||||
dim=-1,
|
||||
output_size=cos_sin.shape[-1])
|
||||
|
||||
query_shape = query.shape
|
||||
query = query.view(num_tokens, -1, self.head_size)
|
||||
query_rot = query[..., :self.rotary_dim]
|
||||
query_pass = query[..., self.rotary_dim:]
|
||||
query_rot = apply_rotary_pos_emb(query_rot, cos, sin, None, 0,
|
||||
rope_mode)
|
||||
query = torch.cat((query_rot, query_pass), dim=-1).reshape(query_shape)
|
||||
|
||||
key_shape = key.shape
|
||||
key = key.view(num_tokens, -1, self.head_size)
|
||||
key_rot = key[..., :self.rotary_dim]
|
||||
key_pass = key[..., self.rotary_dim:]
|
||||
key_rot = apply_rotary_pos_emb(key_rot, cos, sin, None, 0, rope_mode)
|
||||
key = torch.cat((key_rot, key_pass), dim=-1).reshape(key_shape)
|
||||
return query, key
|
||||
|
||||
def extra_repr(self) -> str:
|
||||
s = f"head_size={self.head_size}, rotary_dim={self.rotary_dim}"
|
||||
s += f", max_position_embeddings={self.max_position_embeddings}"
|
||||
s += f", base={self.base}, is_neox_style={self.is_neox_style}"
|
||||
return s
|
||||
|
||||
|
||||
def _to_tuple(x: Union[int, Tuple[int, ...]], dim: int = 2) -> Tuple[int, ...]:
|
||||
|
||||
|
||||
def _to_tuple(x, dim=2):
|
||||
if isinstance(x, int):
|
||||
return (x, ) * dim
|
||||
return (x,) * dim
|
||||
elif len(x) == dim:
|
||||
return x
|
||||
else:
|
||||
raise ValueError(f"Expected length {dim} or int, but got {x}")
|
||||
|
||||
|
||||
def get_meshgrid_nd(start: Union[int, Tuple[int, ...]],
|
||||
*args: Union[int, Tuple[int, ...]],
|
||||
dim: int = 2) -> torch.Tensor:
|
||||
|
||||
|
||||
|
||||
def get_meshgrid_nd(start, *args, dim=2):
|
||||
"""
|
||||
Get n-D meshgrid with start, stop and num.
|
||||
|
||||
@@ -216,7 +296,7 @@ def get_meshgrid_nd(start: Union[int, Tuple[int, ...]],
|
||||
# start is start, args[0] is stop, step is 1
|
||||
start = _to_tuple(start, dim=dim)
|
||||
stop = _to_tuple(args[0], dim=dim)
|
||||
num = tuple(stop[i] - start[i] for i in range(dim))
|
||||
num = [stop[i] - start[i] for i in range(dim)]
|
||||
elif len(args) == 2:
|
||||
# start is start, args[0] is stop, args[1] is num
|
||||
start = _to_tuple(start, dim=dim) # Left-Top eg: 12,0
|
||||
@@ -243,7 +323,6 @@ def get_1d_rotary_pos_embed(
|
||||
theta: float = 10000.0,
|
||||
theta_rescale_factor: float = 1.0,
|
||||
interpolation_factor: float = 1.0,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Precompute the frequency tensor for complex exponential (cis) with given dimensions.
|
||||
@@ -270,14 +349,12 @@ def get_1d_rotary_pos_embed(
|
||||
if theta_rescale_factor != 1.0:
|
||||
theta *= theta_rescale_factor**(dim / (dim - 2))
|
||||
|
||||
freqs = 1.0 / (theta**(torch.arange(0, dim, 2)[:(dim // 2)].to(dtype) / dim)
|
||||
) # [D/2]
|
||||
freqs = 1.0 / (theta**(torch.arange(0, dim, 2)[:(dim // 2)].to(torch.float64) / dim)) # [D/2]
|
||||
freqs = torch.outer(pos * interpolation_factor, freqs) # [S, D/2]
|
||||
freqs_cos = freqs.cos() # [S, D/2]
|
||||
freqs_sin = freqs.sin() # [S, D/2]
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
def get_nd_rotary_pos_embed(
|
||||
rope_dim_list,
|
||||
start,
|
||||
@@ -288,7 +365,6 @@ def get_nd_rotary_pos_embed(
|
||||
shard_dim: int = 0,
|
||||
sp_rank: int = 0,
|
||||
sp_world_size: int = 1,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
This is a n-d version of precompute_freqs_cis, which is a RoPE for tokens with n-d structure.
|
||||
@@ -311,58 +387,51 @@ def get_nd_rotary_pos_embed(
|
||||
Tuple[torch.Tensor, torch.Tensor]: (cos, sin) tensors of shape [HW, D/2]
|
||||
"""
|
||||
# Get the full grid
|
||||
full_grid = get_meshgrid_nd(
|
||||
start, *args, dim=len(rope_dim_list)) # [3, W, H, D] / [2, W, H]
|
||||
|
||||
full_grid = get_meshgrid_nd(start, *args, dim=len(rope_dim_list)) # [3, W, H, D] / [2, W, H]
|
||||
|
||||
# Shard the grid if using sequence parallelism (sp_world_size > 1)
|
||||
assert shard_dim < len(
|
||||
rope_dim_list
|
||||
), f"shard_dim {shard_dim} must be less than number of dimensions {len(rope_dim_list)}"
|
||||
assert shard_dim < len(rope_dim_list), f"shard_dim {shard_dim} must be less than number of dimensions {len(rope_dim_list)}"
|
||||
if sp_world_size > 1:
|
||||
# Get the shape of the full grid
|
||||
grid_shape = list(full_grid.shape[1:])
|
||||
|
||||
|
||||
# Ensure the dimension to shard is divisible by sp_world_size
|
||||
assert grid_shape[shard_dim] % sp_world_size == 0, (
|
||||
f"Dimension {shard_dim} with size {grid_shape[shard_dim]} is not divisible "
|
||||
f"by sequence parallel world size {sp_world_size}")
|
||||
|
||||
f"by sequence parallel world size {sp_world_size}"
|
||||
)
|
||||
|
||||
# Compute the start and end indices for this rank's shard
|
||||
shard_size = grid_shape[shard_dim] // sp_world_size
|
||||
start_idx = sp_rank * shard_size
|
||||
end_idx = (sp_rank + 1) * shard_size
|
||||
|
||||
|
||||
# Create slicing indices for each dimension
|
||||
slice_indices = [slice(None) for _ in range(len(grid_shape))]
|
||||
slice_indices[shard_dim] = slice(start_idx, end_idx)
|
||||
|
||||
|
||||
# Shard the grid
|
||||
# Update grid shape for the sharded dimension
|
||||
grid_shape[shard_dim] = grid_shape[shard_dim] // sp_world_size
|
||||
grid = torch.empty((len(rope_dim_list), ) + tuple(grid_shape),
|
||||
dtype=full_grid.dtype)
|
||||
grid = torch.empty((len(rope_dim_list),) + tuple(grid_shape), dtype=full_grid.dtype)
|
||||
for i in range(len(rope_dim_list)):
|
||||
grid[i] = full_grid[i][tuple(slice_indices)]
|
||||
else:
|
||||
grid = full_grid
|
||||
|
||||
if isinstance(theta_rescale_factor, (int, float)):
|
||||
if isinstance(theta_rescale_factor, int) or isinstance(theta_rescale_factor, float):
|
||||
theta_rescale_factor = [theta_rescale_factor] * len(rope_dim_list)
|
||||
elif isinstance(theta_rescale_factor,
|
||||
list) and len(theta_rescale_factor) == 1:
|
||||
elif isinstance(theta_rescale_factor, list) and len(theta_rescale_factor) == 1:
|
||||
theta_rescale_factor = [theta_rescale_factor[0]] * len(rope_dim_list)
|
||||
assert len(theta_rescale_factor) == len(
|
||||
rope_dim_list
|
||||
), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
|
||||
rope_dim_list), "len(theta_rescale_factor) should equal to len(rope_dim_list)"
|
||||
|
||||
if isinstance(interpolation_factor, (int, float)):
|
||||
if isinstance(interpolation_factor, int) or isinstance(interpolation_factor, float):
|
||||
interpolation_factor = [interpolation_factor] * len(rope_dim_list)
|
||||
elif isinstance(interpolation_factor,
|
||||
list) and len(interpolation_factor) == 1:
|
||||
elif isinstance(interpolation_factor, list) and len(interpolation_factor) == 1:
|
||||
interpolation_factor = [interpolation_factor[0]] * len(rope_dim_list)
|
||||
assert len(interpolation_factor) == len(
|
||||
rope_dim_list
|
||||
), "len(interpolation_factor) should equal to len(rope_dim_list)"
|
||||
rope_dim_list), "len(interpolation_factor) should equal to len(rope_dim_list)"
|
||||
|
||||
# use 1/ndim of dimensions to encode grid_axis
|
||||
embs = []
|
||||
@@ -373,7 +442,6 @@ def get_nd_rotary_pos_embed(
|
||||
theta,
|
||||
theta_rescale_factor=theta_rescale_factor[i],
|
||||
interpolation_factor=interpolation_factor[i],
|
||||
dtype=dtype,
|
||||
) # 2 x [WHD, rope_dim_list[i]]
|
||||
embs.append(emb)
|
||||
|
||||
@@ -383,15 +451,14 @@ def get_nd_rotary_pos_embed(
|
||||
|
||||
|
||||
def get_rotary_pos_embed(
|
||||
rope_sizes,
|
||||
hidden_size,
|
||||
heads_num,
|
||||
rope_dim_list,
|
||||
rope_theta,
|
||||
theta_rescale_factor=1.0,
|
||||
rope_sizes,
|
||||
hidden_size,
|
||||
heads_num,
|
||||
rope_dim_list,
|
||||
rope_theta,
|
||||
theta_rescale_factor=1.0,
|
||||
interpolation_factor=1.0,
|
||||
shard_dim: int = 0,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
"""
|
||||
Generate rotary positional embeddings for the given sizes.
|
||||
@@ -412,19 +479,17 @@ def get_rotary_pos_embed(
|
||||
|
||||
target_ndim = 3
|
||||
head_dim = hidden_size // heads_num
|
||||
|
||||
|
||||
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"
|
||||
|
||||
# Get SP info
|
||||
sp_group = get_sp_group()
|
||||
sp_rank = sp_group.rank_in_group
|
||||
sp_world_size = sp_group.world_size
|
||||
|
||||
|
||||
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
|
||||
rope_dim_list,
|
||||
rope_sizes,
|
||||
@@ -433,15 +498,12 @@ def get_rotary_pos_embed(
|
||||
interpolation_factor=interpolation_factor,
|
||||
shard_dim=shard_dim,
|
||||
sp_rank=sp_rank,
|
||||
sp_world_size=sp_world_size,
|
||||
dtype=dtype,
|
||||
sp_world_size=sp_world_size
|
||||
)
|
||||
return freqs_cos, freqs_sin
|
||||
|
||||
|
||||
_ROPE_DICT: Dict[Tuple, RotaryEmbedding] = {}
|
||||
|
||||
|
||||
def get_rope(
|
||||
head_size: int,
|
||||
rotary_dim: int,
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/layers/utils.py
|
||||
"""Utility methods for model layers."""
|
||||
from typing import Tuple
|
||||
|
||||
@@ -21,3 +20,39 @@ def get_token_bin_counts_and_mask(
|
||||
mask = bin_counts > 0
|
||||
|
||||
return bin_counts, mask
|
||||
|
||||
|
||||
def apply_penalties(logits: torch.Tensor, prompt_tokens_tensor: torch.Tensor,
|
||||
output_tokens_tensor: torch.Tensor,
|
||||
presence_penalties: torch.Tensor,
|
||||
frequency_penalties: torch.Tensor,
|
||||
repetition_penalties: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Applies penalties in place to the logits tensor
|
||||
logits : The input logits tensor of shape [num_seqs, vocab_size]
|
||||
prompt_tokens_tensor: A tensor containing the prompt tokens. The prompts
|
||||
are padded to the maximum prompt length within the batch using
|
||||
`vocab_size` as the padding value. The value `vocab_size` is used
|
||||
for padding because it does not correspond to any valid token ID
|
||||
in the vocabulary.
|
||||
output_tokens_tensor: The output tokens tensor.
|
||||
presence_penalties: The presence penalties of shape (num_seqs, )
|
||||
frequency_penalties: The frequency penalties of shape (num_seqs, )
|
||||
repetition_penalties: The repetition penalties of shape (num_seqs, )
|
||||
"""
|
||||
num_seqs, vocab_size = logits.shape
|
||||
_, prompt_mask = get_token_bin_counts_and_mask(prompt_tokens_tensor,
|
||||
vocab_size, num_seqs)
|
||||
output_bin_counts, output_mask = get_token_bin_counts_and_mask(
|
||||
output_tokens_tensor, vocab_size, num_seqs)
|
||||
repetition_penalties = repetition_penalties.unsqueeze(dim=1).repeat(
|
||||
1, vocab_size)
|
||||
logits[logits > 0] /= torch.where(prompt_mask | output_mask,
|
||||
repetition_penalties, 1.0)[logits > 0]
|
||||
logits[logits <= 0] *= torch.where(prompt_mask | output_mask,
|
||||
repetition_penalties, 1.0)[logits <= 0]
|
||||
# We follow the definition in OpenAI API.
|
||||
# Refer to https://platform.openai.com/docs/api-reference/parameter-details
|
||||
logits -= frequency_penalties.unsqueeze(dim=1) * output_bin_counts
|
||||
logits -= presence_penalties.unsqueeze(dim=1) * output_mask
|
||||
return logits
|
||||
|
||||
@@ -1,16 +1,11 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from typing import Optional
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
import math
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
from typing import Optional
|
||||
from fastvideo.v1.layers.mlp import MLP
|
||||
|
||||
|
||||
class PatchEmbed(nn.Module):
|
||||
"""2D Image to Patch Embedding
|
||||
|
||||
@@ -25,14 +20,16 @@ class PatchEmbed(nn.Module):
|
||||
Remove the _assert function in forward function to be compatible with multi-resolution images.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
patch_size=16,
|
||||
in_chans=3,
|
||||
embed_dim=768,
|
||||
norm_layer=None,
|
||||
flatten=True,
|
||||
bias=True,
|
||||
dtype=None):
|
||||
def __init__(
|
||||
self,
|
||||
patch_size=16,
|
||||
in_chans=3,
|
||||
embed_dim=768,
|
||||
norm_layer=None,
|
||||
flatten=True,
|
||||
bias=True,
|
||||
dtype=None
|
||||
):
|
||||
super().__init__()
|
||||
# Convert patch_size to 2-tuple
|
||||
if isinstance(patch_size, (list, tuple)):
|
||||
@@ -40,16 +37,18 @@ class PatchEmbed(nn.Module):
|
||||
patch_size = (patch_size[0], patch_size[0])
|
||||
else:
|
||||
patch_size = (patch_size, patch_size)
|
||||
|
||||
|
||||
self.patch_size = patch_size
|
||||
self.flatten = flatten
|
||||
|
||||
self.proj = nn.Conv3d(in_chans,
|
||||
embed_dim,
|
||||
kernel_size=patch_size,
|
||||
stride=patch_size,
|
||||
bias=bias,
|
||||
dtype=dtype)
|
||||
self.proj = nn.Conv3d(
|
||||
in_chans,
|
||||
embed_dim,
|
||||
kernel_size=patch_size,
|
||||
stride=patch_size,
|
||||
bias=bias,
|
||||
dtype=dtype
|
||||
)
|
||||
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
|
||||
|
||||
def forward(self, x):
|
||||
@@ -60,11 +59,13 @@ class PatchEmbed(nn.Module):
|
||||
return x
|
||||
|
||||
|
||||
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
"""
|
||||
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size,
|
||||
@@ -72,34 +73,27 @@ class TimestepEmbedder(nn.Module):
|
||||
frequency_embedding_size=256,
|
||||
max_period=10000,
|
||||
dtype=None,
|
||||
freq_dtype=torch.float32,
|
||||
):
|
||||
super().__init__()
|
||||
self.frequency_embedding_size = frequency_embedding_size
|
||||
self.max_period = max_period
|
||||
|
||||
self.mlp = MLP(frequency_embedding_size,
|
||||
hidden_size,
|
||||
hidden_size,
|
||||
act_type=act_layer,
|
||||
dtype=dtype)
|
||||
self.freq_dtype = freq_dtype
|
||||
|
||||
def forward(self, t: torch.Tensor) -> torch.Tensor:
|
||||
t_freq = timestep_embedding(t,
|
||||
self.frequency_embedding_size,
|
||||
self.max_period,
|
||||
dtype=self.freq_dtype).to(
|
||||
self.mlp.fc_in.weight.dtype)
|
||||
|
||||
self.mlp = MLP(
|
||||
frequency_embedding_size,
|
||||
hidden_size,
|
||||
hidden_size,
|
||||
act_type=act_layer,
|
||||
dtype=dtype
|
||||
)
|
||||
|
||||
def forward(self, t):
|
||||
t_freq = timestep_embedding(t, self.frequency_embedding_size, self.max_period).float()
|
||||
# t_freq = t_freq.to(self.mlp.fc_in.weight.dtype)
|
||||
t_emb = self.mlp(t_freq)
|
||||
return t_emb
|
||||
|
||||
|
||||
def timestep_embedding(t: torch.Tensor,
|
||||
dim: int,
|
||||
max_period: int = 10000,
|
||||
dtype: torch.dtype = torch.float32) -> torch.Tensor:
|
||||
def timestep_embedding(t, dim, max_period=10000):
|
||||
"""
|
||||
Create sinusoidal timestep embeddings.
|
||||
|
||||
@@ -112,20 +106,17 @@ def timestep_embedding(t: torch.Tensor,
|
||||
Tensor of shape [B, dim] with embeddings
|
||||
"""
|
||||
half = dim // 2
|
||||
freqs = torch.exp(-math.log(max_period) *
|
||||
torch.arange(start=0, end=half, dtype=dtype) /
|
||||
half).to(device=t.device)
|
||||
freqs = torch.exp(-math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float64) / half).to(device=t.device)
|
||||
args = t[:, None].float() * freqs[None]
|
||||
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
|
||||
if dim % 2:
|
||||
embedding = torch.cat(
|
||||
[embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
|
||||
return embedding
|
||||
|
||||
|
||||
class ModulateProjection(nn.Module):
|
||||
"""Modulation layer for DiT blocks."""
|
||||
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
hidden_size: int,
|
||||
@@ -136,12 +127,14 @@ class ModulateProjection(nn.Module):
|
||||
super().__init__()
|
||||
self.factor = factor
|
||||
self.hidden_size = hidden_size
|
||||
self.linear = ReplicatedLinear(hidden_size,
|
||||
hidden_size * factor,
|
||||
bias=True,
|
||||
params_dtype=dtype)
|
||||
self.linear = ReplicatedLinear(
|
||||
hidden_size,
|
||||
hidden_size * factor,
|
||||
bias=True,
|
||||
params_dtype=dtype
|
||||
)
|
||||
self.act = get_act_fn(act_layer)
|
||||
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
x = self.act(x)
|
||||
x, _ = self.linear(x)
|
||||
@@ -161,13 +154,12 @@ def unpatchify(x, t, h, w, patch_size, channels):
|
||||
"""
|
||||
assert x.ndim == 3, f"x.ndim: {x.ndim}"
|
||||
assert len(patch_size) == 3, f"patch_size: {patch_size}"
|
||||
assert t * h * w == x.shape[
|
||||
1], f"t * h * w: {t * h * w}, x.shape[1]: {x.shape[1]}"
|
||||
assert t * h * w == x.shape[1], f"t * h * w: {t * h * w}, x.shape[1]: {x.shape[1]}"
|
||||
c = channels
|
||||
pt, ph, pw = patch_size
|
||||
|
||||
|
||||
x = x.reshape(shape=(x.shape[0], t, h, w, c, pt, ph, pw))
|
||||
x = torch.einsum("nthwcopq->nctohpwq", x)
|
||||
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
|
||||
|
||||
return imgs
|
||||
|
||||
return imgs
|
||||
@@ -6,12 +6,12 @@ from typing import List, Optional, Sequence, Tuple
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch.nn.parameter import Parameter, UninitializedParameter
|
||||
from vllm.model_executor.layers.quantization.base_config import (
|
||||
QuantizationConfig, QuantizeMethodBase, method_has_implemented_embedding)
|
||||
|
||||
from fastvideo.v1.distributed import (divide, get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size,
|
||||
tensor_model_parallel_all_reduce)
|
||||
get_tensor_model_parallel_world_size,
|
||||
tensor_model_parallel_all_reduce)
|
||||
from vllm.model_executor.layers.quantization.base_config import (
|
||||
QuantizationConfig, QuantizeMethodBase, method_has_implemented_embedding)
|
||||
from fastvideo.v1.models.parameter import BasevLLMParameter
|
||||
from fastvideo.v1.models.utils import set_weight_attrs
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
@@ -53,9 +53,10 @@ def pad_vocab_size(vocab_size: int,
|
||||
return ((vocab_size + pad_to - 1) // pad_to) * pad_to
|
||||
|
||||
|
||||
def vocab_range_from_per_partition_vocab_size(per_partition_vocab_size: int,
|
||||
rank: int,
|
||||
offset: int = 0) -> Sequence[int]:
|
||||
def vocab_range_from_per_partition_vocab_size(
|
||||
per_partition_vocab_size: int,
|
||||
rank: int,
|
||||
offset: int = 0) -> Sequence[int]:
|
||||
index_f = rank * per_partition_vocab_size
|
||||
index_l = index_f + per_partition_vocab_size
|
||||
return index_f + offset, index_l + offset
|
||||
@@ -142,14 +143,14 @@ def get_masked_input_and_mask(
|
||||
added_vocab_end_index: int) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
# torch.compile will fuse all of the pointwise ops below
|
||||
# into a single kernel, making it very fast
|
||||
org_vocab_mask = (input_ >= org_vocab_start_index) & (input_
|
||||
< org_vocab_end_index)
|
||||
org_vocab_mask = (input_ >= org_vocab_start_index) & (
|
||||
input_ < org_vocab_end_index)
|
||||
added_vocab_mask = (input_ >= added_vocab_start_index) & (
|
||||
input_ < added_vocab_end_index)
|
||||
added_offset = added_vocab_start_index - (
|
||||
org_vocab_end_index - org_vocab_start_index) - num_org_vocab_padding
|
||||
valid_offset = (org_vocab_start_index * org_vocab_mask) + (added_offset *
|
||||
added_vocab_mask)
|
||||
valid_offset = (org_vocab_start_index *
|
||||
org_vocab_mask) + (added_offset * added_vocab_mask)
|
||||
vocab_mask = org_vocab_mask | added_vocab_mask
|
||||
input_ = vocab_mask * (input_ - valid_offset)
|
||||
return input_, ~vocab_mask
|
||||
@@ -293,8 +294,8 @@ class VocabParallelEmbedding(torch.nn.Module):
|
||||
return VocabParallelEmbeddingShardIndices(
|
||||
padded_org_vocab_start_index, padded_org_vocab_end_index,
|
||||
padded_added_vocab_start_index, padded_added_vocab_end_index,
|
||||
org_vocab_start_index, org_vocab_end_index, added_vocab_start_index,
|
||||
added_vocab_end_index)
|
||||
org_vocab_start_index, org_vocab_end_index,
|
||||
added_vocab_start_index, added_vocab_end_index)
|
||||
|
||||
def get_sharded_to_full_mapping(self) -> Optional[List[int]]:
|
||||
"""Get a mapping that can be used to reindex the gathered
|
||||
@@ -385,8 +386,19 @@ class VocabParallelEmbedding(torch.nn.Module):
|
||||
# Copy the data. Select chunk corresponding to current shard.
|
||||
loaded_weight = loaded_weight.narrow(output_dim, start_idx, shard_size)
|
||||
|
||||
param[:loaded_weight.shape[0]].data.copy_(loaded_weight)
|
||||
param[loaded_weight.shape[0]:].data.fill_(0)
|
||||
if current_platform.is_hpu():
|
||||
# FIXME(kzawora): Weight copy with slicing bugs out on Gaudi here,
|
||||
# so we're using a workaround. Remove this when fixed in
|
||||
# HPU PT bridge.
|
||||
padded_weight = torch.cat([
|
||||
loaded_weight,
|
||||
torch.zeros(param.shape[0] - loaded_weight.shape[0],
|
||||
*loaded_weight.shape[1:])
|
||||
])
|
||||
param.data.copy_(padded_weight)
|
||||
else:
|
||||
param[:loaded_weight.shape[0]].data.copy_(loaded_weight)
|
||||
param[loaded_weight.shape[0]:].data.fill_(0)
|
||||
|
||||
def forward(self, input_):
|
||||
if self.tp_size > 1:
|
||||
@@ -400,7 +412,8 @@ class VocabParallelEmbedding(torch.nn.Module):
|
||||
else:
|
||||
masked_input = input_
|
||||
# Get the embeddings.
|
||||
output_parallel = self.quant_method.embedding(self, masked_input.long())
|
||||
output_parallel = self.quant_method.embedding(self,
|
||||
masked_input.long())
|
||||
# Mask the output embedding.
|
||||
if self.tp_size > 1:
|
||||
output_parallel.masked_fill_(input_mask.unsqueeze(-1), 0)
|
||||
@@ -415,3 +428,57 @@ class VocabParallelEmbedding(torch.nn.Module):
|
||||
s += f', num_embeddings_padded={self.num_embeddings_padded}'
|
||||
s += f', tp_size={self.tp_size}'
|
||||
return s
|
||||
|
||||
|
||||
class ParallelLMHead(VocabParallelEmbedding):
|
||||
"""Parallelized LM head.
|
||||
|
||||
Output logits weight matrices used in the Sampler. The weight and bias
|
||||
tensors are padded to make sure they are divisible by the number of
|
||||
model parallel GPUs.
|
||||
|
||||
Args:
|
||||
num_embeddings: vocabulary size.
|
||||
embedding_dim: size of hidden state.
|
||||
bias: whether to use bias.
|
||||
params_dtype: type of the parameters.
|
||||
org_num_embeddings: original vocabulary size (without LoRA).
|
||||
padding_size: padding size for the vocabulary.
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
num_embeddings: int,
|
||||
embedding_dim: int,
|
||||
bias: bool = False,
|
||||
params_dtype: Optional[torch.dtype] = None,
|
||||
org_num_embeddings: Optional[int] = None,
|
||||
padding_size: int = DEFAULT_VOCAB_PADDING_SIZE,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__(num_embeddings, embedding_dim, params_dtype,
|
||||
org_num_embeddings, padding_size, quant_config,
|
||||
prefix)
|
||||
self.quant_config = quant_config
|
||||
if bias:
|
||||
self.bias = Parameter(
|
||||
torch.empty(self.num_embeddings_per_partition,
|
||||
dtype=params_dtype))
|
||||
set_weight_attrs(self.bias, {
|
||||
"output_dim": 0,
|
||||
"weight_loader": self.weight_loader,
|
||||
})
|
||||
else:
|
||||
self.register_parameter("bias", None)
|
||||
|
||||
def tie_weights(self, embed_tokens: VocabParallelEmbedding):
|
||||
"""Tie the weights with word embeddings."""
|
||||
# GGUF quantized embed_tokens.
|
||||
if self.quant_config and self.quant_config.get_name() == "gguf":
|
||||
return embed_tokens
|
||||
else:
|
||||
self.weight = embed_tokens.weight
|
||||
return self
|
||||
|
||||
def forward(self, input_):
|
||||
del input_
|
||||
raise RuntimeError("LMHead's weights should be used in the sampler.")
|
||||
|
||||
@@ -1,5 +1,9 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/logger.py
|
||||
# adapted from vllm
|
||||
# https://github.com/vllm-project/vllm/blob/main/vllm/logger.py
|
||||
# Copyright 2023 The vLLM Authors.
|
||||
# Copyright 2023 The FastVideo Authors.
|
||||
|
||||
"""Logging configuration for fastvideo.v1."""
|
||||
import datetime
|
||||
import json
|
||||
@@ -100,8 +104,7 @@ def _configure_fastvideo_root_logger() -> None:
|
||||
"FASTVIDEO_CONFIGURE_LOGGING evaluated to false, but "
|
||||
"FASTVIDEO_LOGGING_CONFIG_PATH was given. FASTVIDEO_LOGGING_CONFIG_PATH "
|
||||
"implies FASTVIDEO_CONFIGURE_LOGGING. Please enable "
|
||||
"FASTVIDEO_CONFIGURE_LOGGING or unset FASTVIDEO_LOGGING_CONFIG_PATH."
|
||||
)
|
||||
"FASTVIDEO_CONFIGURE_LOGGING or unset FASTVIDEO_LOGGING_CONFIG_PATH.")
|
||||
|
||||
if FASTVIDEO_CONFIGURE_LOGGING:
|
||||
logging_config = DEFAULT_LOGGING_CONFIG
|
||||
@@ -127,7 +130,6 @@ def _configure_fastvideo_root_logger() -> None:
|
||||
if logging_config:
|
||||
dictConfig(logging_config)
|
||||
|
||||
|
||||
# TODO: add rank_zero_only log
|
||||
def init_logger(name: str) -> _FastvideoLogger:
|
||||
"""The main purpose of this function is to ensure that loggers are
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
|
||||
from fastvideo.v1.logging_utils.formatter import NewLineFormatter
|
||||
|
||||
__all__ = [
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/logging_utils/formatter.py
|
||||
# adapted from vllm
|
||||
|
||||
import logging
|
||||
|
||||
|
||||
@@ -0,0 +1,31 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch.nn as nn
|
||||
from typing import Dict
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def get_scheduler(module_path: str, architecture: str, inference_args: InferenceArgs) -> Dict:
|
||||
"""Create a scheduler based on the inference args. Can be overridden by subclasses."""
|
||||
if hasattr(inference_args, 'denoise_type') and inference_args.denoise_type == "flow":
|
||||
# TODO(will): add schedulers to register or create a new scheduler registry
|
||||
# TODO(will): default to config file but allow override through
|
||||
# inference args. Currently only uses inference args.
|
||||
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchDiscreteScheduler
|
||||
return FlowMatchDiscreteScheduler(
|
||||
shift=inference_args.flow_shift,
|
||||
solver=inference_args.flow_solver,
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Invalid denoise type: {inference_args.denoise_type}")
|
||||
|
||||
__all__ = [
|
||||
"set_random_seed",
|
||||
"BasevLLMParameter",
|
||||
"PackedvLLMParameter",
|
||||
"get_model",
|
||||
"get_scheduler",
|
||||
]
|
||||
|
||||
@@ -1,15 +1,12 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from torch import nn
|
||||
|
||||
|
||||
# TODO
|
||||
class BaseDiT(nn.Module):
|
||||
_fsdp_shard_conditions: list = []
|
||||
attention_head_dim: int | None = None
|
||||
|
||||
def __init__(self, *args, **kwargs) -> None:
|
||||
_fsdp_shard_conditions = []
|
||||
attention_head_dim: int = None
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__()
|
||||
|
||||
def forward(self, *args, **kwargs):
|
||||
pass
|
||||
pass
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,30 +1,19 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from typing import Any, Dict, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.v1.attention import DistributedAttention, LocalAttention
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, RMSNorm,
|
||||
ScaleResidual,
|
||||
ScaleResidualLayerNormScaleShift)
|
||||
import math
|
||||
from typing import Optional, Tuple, List, Union, Dict, Any
|
||||
from fastvideo.v1.attention.flash_attn import DistributedAttention, LocalAttention
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
from fastvideo.v1.layers.layernorm import LayerNormScaleShift, ScaleResidual, ScaleResidualLayerNormScaleShift, RMSNorm
|
||||
from fastvideo.v1.layers.visual_embedding import PatchEmbed, TimestepEmbedder, ModulateProjection
|
||||
from fastvideo.v1.layers.rotary_embedding import _apply_rotary_emb, get_rotary_pos_embed
|
||||
from fastvideo.v1.distributed.parallel_state import get_sequence_model_parallel_world_size
|
||||
# from torch.nn import RMSNorm
|
||||
# TODO: RMSNorm ....
|
||||
from fastvideo.v1.layers.mlp import MLP
|
||||
from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
|
||||
get_rotary_pos_embed)
|
||||
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
|
||||
PatchEmbed, TimestepEmbedder)
|
||||
from fastvideo.v1.models.dits.base import BaseDiT
|
||||
|
||||
|
||||
class WanImageEmbedding(torch.nn.Module):
|
||||
|
||||
def __init__(self, in_features: int, out_features: int):
|
||||
super().__init__()
|
||||
|
||||
@@ -32,16 +21,13 @@ class WanImageEmbedding(torch.nn.Module):
|
||||
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
|
||||
self.norm2 = nn.LayerNorm(out_features)
|
||||
|
||||
def forward(self,
|
||||
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
hidden_states = self.norm1(encoder_hidden_states_image)
|
||||
hidden_states = self.ff(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
|
||||
class WanTimeTextImageEmbedding(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
@@ -51,19 +37,9 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.time_embedder = TimestepEmbedder(
|
||||
dim,
|
||||
frequency_embedding_size=time_freq_dim,
|
||||
act_layer="silu",
|
||||
freq_dtype=torch.float64)
|
||||
self.time_modulation = ModulateProjection(dim,
|
||||
factor=6,
|
||||
act_layer="silu")
|
||||
self.text_embedder = MLP(text_embed_dim,
|
||||
dim,
|
||||
dim,
|
||||
bias=True,
|
||||
act_type="gelu_pytorch_tanh")
|
||||
self.time_embedder = TimestepEmbedder(dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
|
||||
self.time_modulation = ModulateProjection(dim, factor=6, act_layer="silu")
|
||||
self.text_embedder = MLP(text_embed_dim, dim, dim, bias=True, act_type="gelu_pytorch_tanh")
|
||||
|
||||
self.image_embedder = None
|
||||
if image_embed_dim is not None:
|
||||
@@ -81,22 +57,19 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
|
||||
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
|
||||
if encoder_hidden_states_image is not None:
|
||||
assert self.image_embedder is not None
|
||||
encoder_hidden_states_image = self.image_embedder(
|
||||
encoder_hidden_states_image)
|
||||
encoder_hidden_states_image = self.image_embedder(encoder_hidden_states_image)
|
||||
|
||||
return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image
|
||||
|
||||
|
||||
class WanSelfAttention(nn.Module):
|
||||
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
num_heads: int,
|
||||
dim,
|
||||
num_heads,
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
eps=1e-6,
|
||||
parallel_attention=False) -> None:
|
||||
parallel_attention=False):
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
@@ -116,12 +89,9 @@ class WanSelfAttention(nn.Module):
|
||||
self.norm_k = RMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
|
||||
|
||||
# Scaled dot product attention
|
||||
self.attn = LocalAttention(dropout_rate=0,
|
||||
softmax_scale=None,
|
||||
causal=False)
|
||||
self.attn = LocalAttention(dropout_rate=0, softmax_scale=None, causal=False)
|
||||
|
||||
def forward(self, x: torch.Tensor, context: torch.Tensor,
|
||||
context_lens: int):
|
||||
def forward(self, x, seq_lens, grid_sizes, freqs):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L, num_heads, C / num_heads]
|
||||
@@ -160,11 +130,11 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
class WanI2VCrossAttention(WanSelfAttention):
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
num_heads: int,
|
||||
dim,
|
||||
num_heads,
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
eps=1e-6) -> None:
|
||||
eps=1e-6):
|
||||
super().__init__(dim, num_heads, window_size, qk_norm, eps)
|
||||
|
||||
self.add_k_proj = ReplicatedLinear(dim, dim)
|
||||
@@ -187,8 +157,7 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
|
||||
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
|
||||
v = self.to_v(context)[0].view(b, -1, n, d)
|
||||
k_img = self.norm_added_k.forward_native(
|
||||
self.add_k_proj(context_img)[0]).view(b, -1, n, d)
|
||||
k_img = self.norm_added_k.forward_native(self.add_k_proj(context_img)[0]).view(b, -1, n, d)
|
||||
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
|
||||
img_x = self.attn(q, k_img, v_img)
|
||||
# compute attention
|
||||
@@ -201,9 +170,7 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
x, _ = self.to_out(x)
|
||||
return x
|
||||
|
||||
|
||||
class WanTransformerBlock(nn.Module):
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
@@ -222,10 +189,10 @@ class WanTransformerBlock(nn.Module):
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_out = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.attn1 = DistributedAttention(num_heads=num_heads,
|
||||
head_size=dim // num_heads,
|
||||
dropout_rate=0.0,
|
||||
causal=False)
|
||||
self.attn1 = DistributedAttention(
|
||||
dropout_rate=0.0,
|
||||
causal=False
|
||||
)
|
||||
self.hidden_dim = dim
|
||||
self.num_attention_heads = num_heads
|
||||
dim_head = dim // num_heads
|
||||
@@ -240,32 +207,16 @@ class WanTransformerBlock(nn.Module):
|
||||
print("QK Norm type not supported")
|
||||
raise Exception
|
||||
assert cross_attn_norm is True
|
||||
self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift(
|
||||
dim,
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32)
|
||||
self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift(dim, norm_type="layer", eps=eps, elementwise_affine=True, dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
# I2V
|
||||
self.attn2 = WanI2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
self.attn2 = WanI2VCrossAttention(dim, num_heads, qk_norm=qk_norm, eps=eps)
|
||||
else:
|
||||
# T2V
|
||||
self.attn2 = WanT2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
|
||||
dim,
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
self.attn2 = WanT2VCrossAttention(dim, num_heads, qk_norm=qk_norm, eps=eps)
|
||||
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(dim, norm_type="layer", eps=eps, elementwise_affine=False, dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -278,7 +229,7 @@ class WanTransformerBlock(nn.Module):
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
freqs_cis: Tuple[torch.Tensor, torch.Tensor],
|
||||
freqs_cis: Tuple[torch.Tensor, torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if hidden_states.dim() == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
@@ -287,13 +238,11 @@ class WanTransformerBlock(nn.Module):
|
||||
assert temb.dtype == torch.float32
|
||||
with torch.cuda.amp.autocast(dtype=torch.float32):
|
||||
e = self.scale_shift_table + temb
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(6, dim=1)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = self.norm1(hidden_states.float()).to(
|
||||
dtype=orig_dtype) * (1 + scale_msa) + shift_msa
|
||||
norm_hidden_states = self.norm1(hidden_states.float()).to(dtype=orig_dtype) * (1 + scale_msa) + shift_msa
|
||||
query, _ = self.to_q(norm_hidden_states)
|
||||
key, _ = self.to_k(norm_hidden_states)
|
||||
value, _ = self.to_v(norm_hidden_states)
|
||||
@@ -308,9 +257,7 @@ class WanTransformerBlock(nn.Module):
|
||||
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
|
||||
# Apply rotary embeddings
|
||||
cos, sin = freqs_cis
|
||||
query, key = _apply_rotary_emb(query, cos, sin,
|
||||
is_neox_style=False), _apply_rotary_emb(
|
||||
key, cos, sin, is_neox_style=False)
|
||||
query, key = _apply_rotary_emb(query, cos, sin, is_neox_style=False), _apply_rotary_emb(key, cos, sin, is_neox_style=False)
|
||||
|
||||
attn_output, _ = self.attn1(query, key, value)
|
||||
attn_output = attn_output.flatten(2)
|
||||
@@ -318,15 +265,11 @@ class WanTransformerBlock(nn.Module):
|
||||
attn_output = attn_output.squeeze(1)
|
||||
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
context_lens=None)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
attn_output = self.attn2(norm_hidden_states, context=encoder_hidden_states, context_lens=None)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
@@ -334,54 +277,39 @@ class WanTransformerBlock(nn.Module):
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class WanTransformer3DModel(BaseDiT):
|
||||
_fsdp_shard_conditions = [
|
||||
lambda n, m: "blocks" in n and str.isdigit(n.split(".")[-1]),
|
||||
]
|
||||
_param_names_mapping = {
|
||||
r"^patch_embedding\.(.*)$":
|
||||
r"patch_embedding.proj.\1",
|
||||
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
|
||||
r"condition_embedder.text_embedder.fc_in.\1",
|
||||
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
|
||||
r"condition_embedder.text_embedder.fc_out.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$":
|
||||
r"condition_embedder.time_embedder.mlp.fc_in.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$":
|
||||
r"condition_embedder.time_embedder.mlp.fc_out.\1",
|
||||
r"^condition_embedder\.time_proj\.(.*)$":
|
||||
r"condition_embedder.time_modulation.linear.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
|
||||
r"condition_embedder.image_embedder.ff.fc_in.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$":
|
||||
r"condition_embedder.image_embedder.ff.fc_out.\1",
|
||||
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$":
|
||||
r"blocks.\1.to_q.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$":
|
||||
r"blocks.\1.to_k.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$":
|
||||
r"blocks.\1.to_v.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$":
|
||||
r"blocks.\1.to_out.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
|
||||
r"blocks.\1.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
|
||||
r"blocks.\1.norm_k.\2",
|
||||
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$":
|
||||
r"blocks.\1.attn2.to_out.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$":
|
||||
r"blocks.\1.ffn.fc_in.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
|
||||
r"blocks.\1.ffn.fc_out.\2",
|
||||
r"blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"blocks.\1.self_attn_residual_norm.norm.\2",
|
||||
}
|
||||
r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1",
|
||||
|
||||
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$": r"condition_embedder.text_embedder.fc_in.\1",
|
||||
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$": r"condition_embedder.text_embedder.fc_out.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_in.\1",
|
||||
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$": r"condition_embedder.time_embedder.mlp.fc_out.\1",
|
||||
r"^condition_embedder\.time_proj\.(.*)$": r"condition_embedder.time_modulation.linear.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$": r"condition_embedder.image_embedder.ff.fc_in.\1",
|
||||
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$": r"condition_embedder.image_embedder.ff.fc_out.\1",
|
||||
|
||||
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$": r"blocks.\1.to_q.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$": r"blocks.\1.to_k.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$": r"blocks.\1.to_v.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$": r"blocks.\1.to_out.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2",
|
||||
|
||||
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$": r"blocks.\1.attn2.to_out.\2",
|
||||
|
||||
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
|
||||
|
||||
r"blocks\.(\d+)\.norm2\.(.*)$": r"blocks.\1.self_attn_residual_norm.norm.\2",
|
||||
}
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: Tuple[int, int, int] = (1, 2, 2),
|
||||
text_len=512,
|
||||
patch_size: Tuple[int] = (1, 2, 2),
|
||||
text_len = 512,
|
||||
num_attention_heads: int = 40,
|
||||
attention_head_dim: int = 128,
|
||||
in_channels: int = 16,
|
||||
@@ -391,7 +319,7 @@ class WanTransformer3DModel(BaseDiT):
|
||||
ffn_dim: int = 13824,
|
||||
num_layers: int = 40,
|
||||
cross_attn_norm: bool = True,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
qk_norm: Optional[str] = "rms_norm_across_heads",
|
||||
eps: float = 1e-6,
|
||||
image_dim: Optional[int] = None,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
@@ -407,10 +335,7 @@ class WanTransformer3DModel(BaseDiT):
|
||||
self.text_len = text_len
|
||||
|
||||
# 1. Patch & position embedding
|
||||
self.patch_embedding = PatchEmbed(in_chans=in_channels,
|
||||
embed_dim=inner_dim,
|
||||
patch_size=patch_size,
|
||||
flatten=False)
|
||||
self.patch_embedding = PatchEmbed(in_chans=in_channels, embed_dim=inner_dim, patch_size=patch_size, flatten=False)
|
||||
|
||||
# 2. Condition embeddings
|
||||
self.condition_embedder = WanTimeTextImageEmbedding(
|
||||
@@ -421,25 +346,22 @@ class WanTransformer3DModel(BaseDiT):
|
||||
)
|
||||
|
||||
# 3. Transformer blocks
|
||||
self.blocks = nn.ModuleList([
|
||||
WanTransformerBlock(inner_dim, ffn_dim, num_attention_heads,
|
||||
qk_norm, cross_attn_norm, eps,
|
||||
added_kv_proj_dim) for _ in range(num_layers)
|
||||
])
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
WanTransformerBlock(
|
||||
inner_dim, ffn_dim, num_attention_heads, qk_norm, cross_attn_norm, eps, added_kv_proj_dim
|
||||
)
|
||||
for _ in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 4. Output norm & projection
|
||||
self.norm_out = LayerNormScaleShift(inner_dim,
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(inner_dim,
|
||||
out_channels * math.prod(patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
|
||||
self.norm_out = LayerNormScaleShift(inner_dim, norm_type="layer", eps=eps, elementwise_affine=False, dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size))
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
@@ -463,14 +385,7 @@ class WanTransformer3DModel(BaseDiT):
|
||||
# Get rotary embeddings
|
||||
d = self.inner_dim // self.num_attention_heads
|
||||
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
|
||||
freqs_cos, freqs_sin = get_rotary_pos_embed(
|
||||
(post_patch_num_frames * get_sequence_model_parallel_world_size(),
|
||||
post_patch_height, post_patch_width),
|
||||
self.inner_dim,
|
||||
self.num_attention_heads,
|
||||
rope_dim_list,
|
||||
dtype=torch.float64,
|
||||
rope_theta=10000)
|
||||
freqs_cos, freqs_sin = get_rotary_pos_embed((post_patch_num_frames * get_sequence_model_parallel_world_size(), post_patch_height, post_patch_width), self.inner_dim, self.num_attention_heads, rope_dim_list, rope_theta=10000)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
|
||||
@@ -481,52 +396,39 @@ class WanTransformer3DModel(BaseDiT):
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
if seq_len is None:
|
||||
seq_len = hidden_states.size(1)
|
||||
hidden_states = torch.cat([
|
||||
hidden_states,
|
||||
hidden_states.new_zeros(1, seq_len - hidden_states.size(1),
|
||||
hidden_states.size(2))
|
||||
],
|
||||
dim=1)
|
||||
hidden_states = torch.cat([hidden_states, hidden_states.new_zeros(1, seq_len - hidden_states.size(1), hidden_states.size(2))], dim=1)
|
||||
|
||||
encoder_hidden_states = torch.cat([
|
||||
encoder_hidden_states,
|
||||
encoder_hidden_states.new_zeros(
|
||||
1, self.text_len - encoder_hidden_states.size(1),
|
||||
encoder_hidden_states.size(2))
|
||||
],
|
||||
dim=1)
|
||||
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image)
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image
|
||||
)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, -1))
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
encoder_hidden_states = torch.concat(
|
||||
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block, hidden_states, encoder_hidden_states, timestep_proj,
|
||||
freqs_cis)
|
||||
block, hidden_states, encoder_hidden_states, timestep_proj, freqs_cis
|
||||
)
|
||||
else:
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis)
|
||||
|
||||
hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, freqs_cis)
|
||||
|
||||
# 5. Output norm, projection & unpatchify
|
||||
with torch.cuda.amp.autocast(dtype=torch.float32):
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(
|
||||
2, dim=1)
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
|
||||
hidden_states = self.norm_out(hidden_states.float(), shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
output = self.unpatchify(hidden_states, grid_sizes)
|
||||
|
||||
return output.float()
|
||||
|
||||
def unpatchify(self, x, grid_sizes) -> torch.Tensor:
|
||||
|
||||
def unpatchify(self, x, grid_sizes):
|
||||
r"""
|
||||
Reconstruct video tensors from patch embeddings.
|
||||
|
||||
@@ -551,4 +453,4 @@ class WanTransformer3DModel(BaseDiT):
|
||||
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
|
||||
out.append(u)
|
||||
out = torch.cat(out, dim=0)
|
||||
return out
|
||||
return out
|
||||
@@ -1,37 +1,35 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/models/clip.py
|
||||
# Adapted from transformers: https://github.com/huggingface/transformers/blob/v4.39.0/src/transformers/models/clip/modeling_clip.py
|
||||
"""Minimal implementation of CLIPVisionModel intended to be only used
|
||||
within a vision language model."""
|
||||
from typing import Iterable, Optional, Set, Tuple, Union, cast
|
||||
from typing import Iterable, Optional, Set, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from transformers import CLIPTextConfig, CLIPVisionConfig
|
||||
from transformers import CLIPVisionConfig, CLIPTextConfig
|
||||
from transformers.modeling_outputs import BaseModelOutputWithPooling
|
||||
from vllm.model_executor.models.interfaces import SupportsQuant
|
||||
|
||||
# from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask, _prepare_4d_attention_mask
|
||||
from fastvideo.v1.attention import LocalAttention
|
||||
from fastvideo.v1.distributed import (divide,
|
||||
get_tensor_model_parallel_world_size)
|
||||
|
||||
from vllm.attention.layer import MultiHeadAttention
|
||||
# from fastvideo.v1.attention.flash_attn import LocalAttention
|
||||
from fastvideo.v1.distributed import divide, get_tensor_model_parallel_world_size
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
|
||||
RowParallelLinear)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.encoders.vision import (VisionEncoderInfo,
|
||||
resolve_visual_encoder_outputs)
|
||||
from fastvideo.v1.layers.linear import (ColumnParallelLinear,
|
||||
QKVParallelLinear,
|
||||
RowParallelLinear)
|
||||
# TODO: support quantization
|
||||
# from vllm.model_executor.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
|
||||
from vllm.model_executor.models.interfaces import SupportsQuant
|
||||
|
||||
from .vision import VisionEncoderInfo, resolve_visual_encoder_outputs
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class QuantizationConfig:
|
||||
pass
|
||||
|
||||
|
||||
class CLIPEncoderInfo(VisionEncoderInfo[CLIPVisionConfig]):
|
||||
|
||||
def get_num_image_tokens(
|
||||
@@ -46,10 +44,10 @@ class CLIPEncoderInfo(VisionEncoderInfo[CLIPVisionConfig]):
|
||||
return self.get_patch_grid_length()**2 + 1
|
||||
|
||||
def get_image_size(self) -> int:
|
||||
return cast(int, self.vision_config.image_size)
|
||||
return self.vision_config.image_size
|
||||
|
||||
def get_patch_size(self) -> int:
|
||||
return cast(int, self.vision_config.patch_size)
|
||||
return self.vision_config.patch_size
|
||||
|
||||
def get_patch_grid_length(self) -> int:
|
||||
image_size, patch_size = self.get_image_size(), self.get_patch_size()
|
||||
@@ -108,14 +106,12 @@ class CLIPTextEmbeddings(nn.Module):
|
||||
embed_dim = config.hidden_size
|
||||
|
||||
self.token_embedding = nn.Embedding(config.vocab_size, embed_dim)
|
||||
self.position_embedding = nn.Embedding(config.max_position_embeddings,
|
||||
embed_dim)
|
||||
self.position_embedding = nn.Embedding(config.max_position_embeddings, embed_dim)
|
||||
|
||||
# position_ids (1, len position emb) is contiguous in memory and exported when serialized
|
||||
self.register_buffer(
|
||||
"position_ids",
|
||||
torch.arange(config.max_position_embeddings).expand((1, -1)),
|
||||
persistent=False)
|
||||
"position_ids", torch.arange(config.max_position_embeddings).expand((1, -1)), persistent=False
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -123,14 +119,7 @@ class CLIPTextEmbeddings(nn.Module):
|
||||
position_ids: Optional[torch.LongTensor] = None,
|
||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if input_ids is not None:
|
||||
seq_length = input_ids.shape[-1]
|
||||
elif inputs_embeds is not None:
|
||||
seq_length = inputs_embeds.shape[-2]
|
||||
else:
|
||||
raise ValueError(
|
||||
"Either input_ids or inputs_embeds must be provided.")
|
||||
|
||||
seq_length = input_ids.shape[-1] if input_ids is not None else inputs_embeds.shape[-2]
|
||||
max_position_embedding = self.position_embedding.weight.shape[0]
|
||||
|
||||
if seq_length > max_position_embedding:
|
||||
@@ -191,11 +180,8 @@ class CLIPAttention(nn.Module):
|
||||
self.tp_size = get_tensor_model_parallel_world_size()
|
||||
self.num_heads_per_partition = divide(self.num_heads, self.tp_size)
|
||||
|
||||
self.attn = LocalAttention(self.num_heads_per_partition,
|
||||
self.head_dim,
|
||||
self.num_heads_per_partition,
|
||||
softmax_scale=self.scale,
|
||||
causal=True)
|
||||
self.attn = MultiHeadAttention(self.num_heads_per_partition,
|
||||
self.head_dim, self.scale)
|
||||
|
||||
def _shape(self, tensor: torch.Tensor, seq_len: int, bsz: int):
|
||||
return tensor.view(bsz, seq_len, self.num_heads,
|
||||
@@ -210,22 +196,12 @@ class CLIPAttention(nn.Module):
|
||||
qkv_states, _ = self.qkv_proj(hidden_states)
|
||||
query_states, key_states, value_states = qkv_states.chunk(3, dim=-1)
|
||||
# use flash_attn_func
|
||||
query_states = query_states.reshape(query_states.shape[0],
|
||||
query_states.shape[1],
|
||||
self.num_heads_per_partition,
|
||||
self.head_dim)
|
||||
key_states = key_states.reshape(key_states.shape[0],
|
||||
key_states.shape[1],
|
||||
self.num_heads_per_partition,
|
||||
self.head_dim)
|
||||
value_states = value_states.reshape(value_states.shape[0],
|
||||
value_states.shape[1],
|
||||
self.num_heads_per_partition,
|
||||
self.head_dim)
|
||||
attn_output = self.attn(query_states, key_states, value_states)
|
||||
attn_output = attn_output.reshape(
|
||||
attn_output.shape[0], attn_output.shape[1],
|
||||
self.num_heads_per_partition * self.head_dim)
|
||||
from flash_attn import flash_attn_func
|
||||
query_states = query_states.reshape(query_states.shape[0], query_states.shape[1], self.num_heads_per_partition, self.head_dim)
|
||||
key_states = key_states.reshape(key_states.shape[0], key_states.shape[1], self.num_heads_per_partition, self.head_dim)
|
||||
value_states = value_states.reshape(value_states.shape[0], value_states.shape[1], self.num_heads_per_partition, self.head_dim)
|
||||
attn_output = flash_attn_func(query_states, key_states, value_states,softmax_scale=self.scale, causal=True)
|
||||
attn_output = attn_output.reshape(attn_output.shape[0], attn_output.shape[1], self.num_heads_per_partition * self.head_dim)
|
||||
attn_output, _ = self.out_proj(attn_output)
|
||||
|
||||
return attn_output, None
|
||||
@@ -345,11 +321,12 @@ class CLIPEncoder(nn.Module):
|
||||
if return_all_hidden_states:
|
||||
return hidden_states_pool
|
||||
return [hidden_states]
|
||||
|
||||
|
||||
|
||||
class CLIPTextTransformer(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
def __init__(self,
|
||||
config: CLIPTextConfig,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
*,
|
||||
@@ -361,14 +338,12 @@ class CLIPTextTransformer(nn.Module):
|
||||
|
||||
self.embeddings = CLIPTextEmbeddings(config)
|
||||
|
||||
self.encoder = CLIPEncoder(
|
||||
config,
|
||||
quant_config=quant_config,
|
||||
num_hidden_layers_override=num_hidden_layers_override,
|
||||
prefix=prefix)
|
||||
self.encoder = CLIPEncoder(config,
|
||||
quant_config=quant_config,
|
||||
num_hidden_layers_override=num_hidden_layers_override,
|
||||
prefix=prefix)
|
||||
|
||||
self.final_layer_norm = nn.LayerNorm(embed_dim,
|
||||
eps=config.layer_norm_eps)
|
||||
self.final_layer_norm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)
|
||||
|
||||
# For `pooled_output` computation
|
||||
self.eos_token_id = config.eos_token_id
|
||||
@@ -390,9 +365,9 @@ class CLIPTextTransformer(nn.Module):
|
||||
|
||||
"""
|
||||
output_attentions = output_attentions if output_attentions is not None else self.config.output_attentions
|
||||
output_hidden_states = (output_hidden_states
|
||||
if output_hidden_states is not None else
|
||||
self.config.output_hidden_states)
|
||||
output_hidden_states = (
|
||||
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
|
||||
)
|
||||
return_dict = return_dict if return_dict is not None else self.config.use_return_dict
|
||||
|
||||
if input_ids is None:
|
||||
@@ -401,8 +376,7 @@ class CLIPTextTransformer(nn.Module):
|
||||
input_shape = input_ids.size()
|
||||
input_ids = input_ids.view(-1, input_shape[-1])
|
||||
|
||||
hidden_states = self.embeddings(input_ids=input_ids,
|
||||
position_ids=position_ids)
|
||||
hidden_states = self.embeddings(input_ids=input_ids, position_ids=position_ids)
|
||||
|
||||
# CLIP's text model uses causal mask, prepare it here.
|
||||
# https://github.com/openai/CLIP/blob/cfcffb90e69f37bf2ff1e988237a0fbe41f33c04/clip/model.py#L324
|
||||
@@ -436,26 +410,24 @@ class CLIPTextTransformer(nn.Module):
|
||||
# take features from the eot embedding (eot_token is the highest number in each sequence)
|
||||
# casting to torch.int for onnx compatibility: argmax doesn't support int64 inputs with opset 14
|
||||
pooled_output = last_hidden_state[
|
||||
torch.arange(last_hidden_state.shape[0],
|
||||
device=last_hidden_state.device),
|
||||
input_ids.to(dtype=torch.int, device=last_hidden_state.device).
|
||||
argmax(dim=-1),
|
||||
torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device),
|
||||
input_ids.to(dtype=torch.int, device=last_hidden_state.device).argmax(dim=-1),
|
||||
]
|
||||
else:
|
||||
# The config gets updated `eos_token_id` from PR #24773 (so the use of exta new tokens is possible)
|
||||
pooled_output = last_hidden_state[
|
||||
torch.arange(last_hidden_state.shape[0],
|
||||
device=last_hidden_state.device),
|
||||
torch.arange(last_hidden_state.shape[0], device=last_hidden_state.device),
|
||||
# We need to get the first position of `eos_token_id` value (`pad_token_ids` might equal to `eos_token_id`)
|
||||
# Note: we assume each sequence (along batch dim.) contains an `eos_token_id` (e.g. prepared by the tokenizer)
|
||||
(input_ids.to(dtype=torch.int, device=last_hidden_state.device
|
||||
) == self.eos_token_id).int().argmax(dim=-1),
|
||||
(input_ids.to(dtype=torch.int, device=last_hidden_state.device) == self.eos_token_id)
|
||||
.int()
|
||||
.argmax(dim=-1),
|
||||
]
|
||||
|
||||
if not return_dict:
|
||||
return (last_hidden_state, pooled_output) + encoder_outputs[1:]
|
||||
|
||||
# return last_hidden_state
|
||||
# return last_hidden_state
|
||||
return BaseModelOutputWithPooling(
|
||||
last_hidden_state=last_hidden_state,
|
||||
pooler_output=pooled_output,
|
||||
@@ -475,10 +447,11 @@ class CLIPTextModel(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
self.config = config
|
||||
self.text_model = CLIPTextTransformer(config=config,
|
||||
quant_config=quant_config,
|
||||
prefix=prefix)
|
||||
|
||||
self.text_model = CLIPTextTransformer(
|
||||
config=config,
|
||||
quant_config=quant_config,
|
||||
prefix=prefix)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.Tensor] = None,
|
||||
@@ -500,8 +473,8 @@ class CLIPTextModel(nn.Module):
|
||||
)
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
|
||||
# Define mapping for stacked parameters
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
@@ -517,7 +490,7 @@ class CLIPTextModel(nn.Module):
|
||||
if weight_name in name:
|
||||
# Replace the weight name with the parameter name
|
||||
model_param_name = name.replace(weight_name, param_name)
|
||||
|
||||
|
||||
if model_param_name in params_dict:
|
||||
param = params_dict[model_param_name]
|
||||
weight_loader = param.weight_loader
|
||||
@@ -528,11 +501,10 @@ class CLIPTextModel(nn.Module):
|
||||
# Use default weight loader for all other parameters
|
||||
if name in params_dict:
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader",
|
||||
default_weight_loader)
|
||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(name)
|
||||
|
||||
|
||||
return loaded_params
|
||||
|
||||
|
||||
@@ -569,7 +541,8 @@ class CLIPVisionTransformer(nn.Module):
|
||||
if len(self.encoder.layers) > config.num_hidden_layers:
|
||||
raise ValueError(
|
||||
f"The original encoder only has {num_hidden_layers} "
|
||||
f"layers, but you requested {len(self.encoder.layers)} layers.")
|
||||
f"layers, but you requested {len(self.encoder.layers)} layers."
|
||||
)
|
||||
|
||||
# If possible, skip post_layernorm to conserve memory
|
||||
if require_post_norm is None:
|
||||
@@ -680,4 +653,4 @@ class CLIPVisionModel(nn.Module, SupportsQuant):
|
||||
default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(name)
|
||||
return loaded_params
|
||||
return loaded_params
|
||||
@@ -1,5 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/models/llama.py
|
||||
|
||||
# Adapted from
|
||||
# https://github.com/huggingface/transformers/blob/v4.28.0/src/transformers/models/llama/modeling_llama.py
|
||||
@@ -23,27 +22,29 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
"""Inference-only LLaMA model compatible with HuggingFace weights."""
|
||||
from typing import Any, Dict, Iterable, Optional, Set, Tuple, Type
|
||||
from typing import Any, Dict, Iterable, Optional, Set, Tuple, Type, Union
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from transformers import LlamaConfig
|
||||
from transformers.modeling_outputs import BaseModelOutputWithPast
|
||||
|
||||
# from vllm.model_executor.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.attention import LocalAttention
|
||||
from vllm.attention.layer import MultiHeadAttention
|
||||
|
||||
from fastvideo.v1.distributed import get_tensor_model_parallel_world_size
|
||||
from fastvideo.v1.layers.activation import SiluAndMul
|
||||
from fastvideo.v1.layers.layernorm import RMSNorm
|
||||
from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
|
||||
QKVParallelLinear, RowParallelLinear)
|
||||
QKVParallelLinear,
|
||||
RowParallelLinear)
|
||||
# from vllm.model_executor.layers.quantization import QuantizationConfig
|
||||
|
||||
from fastvideo.v1.layers.rotary_embedding import get_rope
|
||||
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.v1.models.loader.weight_utils import (default_weight_loader,
|
||||
maybe_remap_kv_scale_name)
|
||||
|
||||
# from ..utils import (extract_layer_index)
|
||||
from fastvideo.v1.models.loader.weight_utils import (
|
||||
default_weight_loader, maybe_remap_kv_scale_name)
|
||||
|
||||
from .utils import (extract_layer_index)
|
||||
|
||||
class QuantizationConfig:
|
||||
pass
|
||||
@@ -103,7 +104,7 @@ class LlamaAttention(nn.Module):
|
||||
bias_o_proj: bool = False,
|
||||
prefix: str = "") -> None:
|
||||
super().__init__()
|
||||
# layer_idx = extract_layer_index(prefix)
|
||||
layer_idx = extract_layer_index(prefix)
|
||||
self.hidden_size = hidden_size
|
||||
tp_size = get_tensor_model_parallel_world_size()
|
||||
self.total_num_heads = num_heads
|
||||
@@ -150,8 +151,7 @@ class LlamaAttention(nn.Module):
|
||||
)
|
||||
|
||||
is_neox_style = True
|
||||
is_gguf = quant_config and hasattr(
|
||||
quant_config, "get_name") and quant_config.get_name() == "gguf"
|
||||
is_gguf = quant_config and quant_config.get_name() == "gguf"
|
||||
if is_gguf and config.model_type == "llama":
|
||||
is_neox_style = False
|
||||
|
||||
@@ -164,11 +164,10 @@ class LlamaAttention(nn.Module):
|
||||
is_neox_style=is_neox_style,
|
||||
)
|
||||
|
||||
self.attn = LocalAttention(self.num_heads,
|
||||
self.head_dim,
|
||||
self.num_kv_heads,
|
||||
softmax_scale=self.scaling,
|
||||
causal=True)
|
||||
self.attn = MultiHeadAttention(self.num_heads,
|
||||
self.head_dim,
|
||||
self.scaling,
|
||||
self.num_kv_heads)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
@@ -181,7 +180,7 @@ class LlamaAttention(nn.Module):
|
||||
# attn_output = self.attn(q, k, v)
|
||||
# use flash_attn_func
|
||||
# TODO (Attn abstraction and backend)
|
||||
# from flash_attn import flash_attn_func
|
||||
from flash_attn import flash_attn_func
|
||||
# reshape q, k, v to (batch_size, seq_len, num_heads, head_dim)
|
||||
batch_size = q.shape[0]
|
||||
seq_len = q.shape[1]
|
||||
@@ -189,10 +188,8 @@ class LlamaAttention(nn.Module):
|
||||
k = k.reshape(batch_size, seq_len, self.num_kv_heads, self.head_dim)
|
||||
v = v.reshape(batch_size, seq_len, self.num_kv_heads, self.head_dim)
|
||||
# import pdb; pdb.set_trace()
|
||||
# attn_output = flash_attn_func(q, k, v, softmax_scale=self.scaling, causal=True)
|
||||
attn_output = self.attn(q, k, v)
|
||||
attn_output = attn_output.reshape(batch_size, seq_len,
|
||||
self.num_heads * self.head_dim)
|
||||
attn_output = flash_attn_func(q, k, v, softmax_scale=self.scaling, causal=True)
|
||||
attn_output = attn_output.reshape(batch_size, seq_len, self.num_heads * self.head_dim)
|
||||
|
||||
output, _ = self.o_proj(attn_output)
|
||||
return output
|
||||
@@ -268,6 +265,7 @@ class LlamaDecoderLayer(nn.Module):
|
||||
|
||||
hidden_states = self.self_attn(positions=positions,
|
||||
hidden_states=hidden_states)
|
||||
|
||||
|
||||
# Fully Connected
|
||||
hidden_states, residual = self.post_attention_layernorm(
|
||||
@@ -289,33 +287,26 @@ class LlamaModel(nn.Module):
|
||||
|
||||
self.config = config
|
||||
self.quant_config = quant_config
|
||||
if lora_config is not None:
|
||||
max_loras = 1
|
||||
lora_vocab_size = 1
|
||||
if hasattr(lora_config, "max_loras"):
|
||||
max_loras = lora_config.max_loras
|
||||
if hasattr(lora_config, "lora_extra_vocab_size"):
|
||||
lora_vocab_size = lora_config.lora_extra_vocab_size
|
||||
lora_vocab = lora_vocab_size * max_loras
|
||||
else:
|
||||
lora_vocab = 0
|
||||
lora_vocab = (lora_config.lora_extra_vocab_size *
|
||||
(lora_config.max_loras or 1)) if lora_config else 0
|
||||
self.vocab_size = config.vocab_size + lora_vocab
|
||||
self.org_vocab_size = config.vocab_size
|
||||
|
||||
|
||||
self.embed_tokens = VocabParallelEmbedding(
|
||||
self.vocab_size,
|
||||
config.hidden_size,
|
||||
org_num_embeddings=config.vocab_size,
|
||||
quant_config=quant_config,
|
||||
)
|
||||
|
||||
|
||||
self.layers = nn.ModuleList([
|
||||
layer_type(config=config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.layers.{i}")
|
||||
for i in range(config.num_hidden_layers)
|
||||
layer_type(
|
||||
config=config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.layers.{i}"
|
||||
) for i in range(config.num_hidden_layers)
|
||||
])
|
||||
|
||||
|
||||
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
|
||||
|
||||
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
|
||||
@@ -329,9 +320,9 @@ class LlamaModel(nn.Module):
|
||||
inputs_embeds: Optional[torch.Tensor] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
) -> torch.Tensor:
|
||||
output_hidden_states = (output_hidden_states
|
||||
if output_hidden_states is not None else
|
||||
self.config.output_hidden_states)
|
||||
output_hidden_states = (
|
||||
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
|
||||
)
|
||||
if inputs_embeds is not None:
|
||||
hidden_states = inputs_embeds
|
||||
else:
|
||||
@@ -339,26 +330,22 @@ class LlamaModel(nn.Module):
|
||||
residual = None
|
||||
|
||||
if positions is None:
|
||||
positions = torch.arange(0,
|
||||
hidden_states.shape[1],
|
||||
device=hidden_states.device).unsqueeze(0)
|
||||
positions = torch.arange(
|
||||
0, hidden_states.shape[1], device=hidden_states.device
|
||||
).unsqueeze(0)
|
||||
|
||||
all_hidden_states: Optional[Tuple[Any, ...]] = (
|
||||
) if output_hidden_states else None
|
||||
all_hidden_states = () if output_hidden_states else None
|
||||
for layer in self.layers:
|
||||
if all_hidden_states is not None:
|
||||
# TODO
|
||||
all_hidden_states += (
|
||||
hidden_states, ) if residual is None else (hidden_states +
|
||||
residual, )
|
||||
if output_hidden_states:
|
||||
all_hidden_states += (hidden_states,)
|
||||
hidden_states, residual = layer(positions, hidden_states, residual)
|
||||
|
||||
hidden_states, _ = self.norm(hidden_states, residual)
|
||||
|
||||
# add hidden states from the last decoder layer
|
||||
if all_hidden_states is not None:
|
||||
all_hidden_states += (hidden_states, )
|
||||
|
||||
if output_hidden_states:
|
||||
all_hidden_states += (hidden_states,)
|
||||
|
||||
# TODO(will): maybe unify the output format with other models and use
|
||||
# our own class
|
||||
output = BaseModelOutputWithPast(
|
||||
@@ -403,12 +390,9 @@ class LlamaModel(nn.Module):
|
||||
continue
|
||||
if "scale" in name:
|
||||
# Remapping the name of FP8 kv-scale.
|
||||
kv_scale_name: Optional[str] = maybe_remap_kv_scale_name(
|
||||
name, params_dict)
|
||||
if kv_scale_name is None:
|
||||
name = maybe_remap_kv_scale_name(name, params_dict)
|
||||
if name is None:
|
||||
continue
|
||||
else:
|
||||
name = kv_scale_name
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
|
||||
+424
-304
@@ -1,6 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from transformers: https://github.com/huggingface/transformers/blob/v4.39.0/src/transformers/models/t5/modeling_t5.py
|
||||
|
||||
# Derived from T5 implementation posted on HuggingFace; license below:
|
||||
#
|
||||
# coding=utf-8
|
||||
@@ -17,48 +15,69 @@
|
||||
# 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.
|
||||
"""PyTorch T5 & UMT5 model."""
|
||||
"""PyTorch T5 model."""
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Iterable, Optional, Set, Tuple
|
||||
import re
|
||||
from typing import Iterable, List, Optional, Set, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from torch import nn
|
||||
from transformers import T5Config
|
||||
# TODO best way to handle xformers imports?
|
||||
from xformers.ops.fmha.attn_bias import LowerTriangularMaskWithTensorBias
|
||||
|
||||
from fastvideo.v1.distributed import get_tensor_model_parallel_world_size
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
from fastvideo.v1.layers.layernorm import RMSNorm
|
||||
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
|
||||
RowParallelLinear)
|
||||
from fastvideo.v1.layers.vocab_parallel_embedding import VocabParallelEmbedding
|
||||
from fastvideo.v1.models.loader.weight_utils import default_weight_loader
|
||||
# TODO func should be in backend interface
|
||||
from vllm.attention.backends.xformers import (XFormersMetadata, _get_attn_bias,
|
||||
_set_attn_bias)
|
||||
from vllm.attention.layer import Attention, AttentionMetadata, AttentionType
|
||||
from vllm.config import CacheConfig, VllmConfig
|
||||
from vllm.distributed import get_tensor_model_parallel_world_size
|
||||
from vllm.model_executor.layers.activation import get_act_fn
|
||||
from vllm.model_executor.layers.linear import (ColumnParallelLinear,
|
||||
QKVParallelLinear,
|
||||
RowParallelLinear)
|
||||
from vllm.model_executor.layers.logits_processor import LogitsProcessor
|
||||
from vllm.model_executor.layers.quantization.base_config import (
|
||||
QuantizationConfig)
|
||||
from vllm.model_executor.layers.sampler import SamplerOutput, get_sampler
|
||||
from vllm.model_executor.layers.vocab_parallel_embedding import (
|
||||
ParallelLMHead, VocabParallelEmbedding)
|
||||
from vllm.model_executor.model_loader.weight_utils import default_weight_loader
|
||||
from vllm.model_executor.sampling_metadata import SamplingMetadata
|
||||
from vllm.sequence import IntermediateTensors
|
||||
|
||||
from .utils import maybe_prefix
|
||||
|
||||
|
||||
class QuantizationConfig:
|
||||
pass
|
||||
class T5LayerNorm(nn.Module):
|
||||
|
||||
def __init__(self, hidden_size, eps=1e-6):
|
||||
"""
|
||||
Construct a layernorm module in the T5 style.
|
||||
No bias and no subtraction of mean.
|
||||
"""
|
||||
super().__init__()
|
||||
self.weight = nn.Parameter(torch.ones(hidden_size))
|
||||
self.variance_epsilon = eps
|
||||
|
||||
class AttentionType:
|
||||
"""
|
||||
Attention type.
|
||||
Use string to be compatible with `torch.compile`.
|
||||
"""
|
||||
# Decoder attention between previous layer Q/K/V
|
||||
DECODER = "decoder"
|
||||
# Encoder attention between previous layer Q/K/V for encoder-decoder
|
||||
ENCODER = "encoder"
|
||||
# Encoder attention between previous layer Q/K/V
|
||||
ENCODER_ONLY = "encoder_only"
|
||||
# Attention between dec. Q and enc. K/V for encoder-decoder
|
||||
ENCODER_DECODER = "encoder_decoder"
|
||||
def forward(self, hidden_states) -> torch.Tensor:
|
||||
# T5 uses a layer_norm which only scales and doesn't shift, which is
|
||||
# also known as Root Mean Square Layer Normalization
|
||||
# https://arxiv.org/abs/1910.07467 thus variance is calculated w/o mean
|
||||
# and there is no bias. Additionally we want to make sure that the
|
||||
# accumulation for half-precision inputs is done in fp32.
|
||||
# TODO (rmns norm ops)
|
||||
variance = hidden_states.to(torch.float32).pow(2).mean(-1,
|
||||
keepdim=True)
|
||||
hidden_states = hidden_states * torch.rsqrt(variance +
|
||||
self.variance_epsilon)
|
||||
|
||||
# convert into half-precision if necessary
|
||||
if self.weight.dtype in [torch.float16, torch.bfloat16]:
|
||||
hidden_states = hidden_states.to(self.weight.dtype)
|
||||
|
||||
@dataclass
|
||||
class AttentionMetadata:
|
||||
attn_bias: torch.Tensor
|
||||
return self.weight * hidden_states
|
||||
|
||||
|
||||
class T5DenseActDense(nn.Module):
|
||||
@@ -77,6 +96,12 @@ class T5DenseActDense(nn.Module):
|
||||
def forward(self, hidden_states) -> torch.Tensor:
|
||||
hidden_states, _ = self.wi(hidden_states)
|
||||
hidden_states = self.act(hidden_states)
|
||||
# if (
|
||||
# isinstance(self.wo.weight, torch.Tensor)
|
||||
# and hidden_states.dtype != self.wo.weight.dtype
|
||||
# and self.wo.weight.dtype != torch.int8
|
||||
# ):
|
||||
# hidden_states = hidden_states.to(self.wo.weight.dtype)
|
||||
hidden_states, _ = self.wo(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
@@ -124,39 +149,23 @@ class T5LayerFF(nn.Module):
|
||||
self.DenseReluDense = T5DenseActDense(config,
|
||||
quant_config=quant_config)
|
||||
|
||||
self.layer_norm = RMSNorm(config.d_model, eps=config.layer_norm_epsilon)
|
||||
self.layer_norm = T5LayerNorm(config.d_model,
|
||||
eps=config.layer_norm_epsilon)
|
||||
|
||||
def forward(self, hidden_states) -> torch.Tensor:
|
||||
forwarded_states = self.layer_norm.forward_native(hidden_states)
|
||||
forwarded_states = self.layer_norm(hidden_states)
|
||||
forwarded_states = self.DenseReluDense(forwarded_states)
|
||||
hidden_states = hidden_states + forwarded_states
|
||||
return hidden_states
|
||||
|
||||
|
||||
# T5 has attn_bias and does not use softmax scaling
|
||||
class T5MultiHeadAttention(nn.Module):
|
||||
|
||||
def __init__(self) -> None:
|
||||
super().__init__()
|
||||
|
||||
def forward(self, q, k, v, attn_bias=None):
|
||||
b, _, n, c = q.shape
|
||||
attn = torch.einsum('binc,bjnc->bnij', q, k)
|
||||
if attn_bias is not None:
|
||||
attn += attn_bias
|
||||
|
||||
attn = F.softmax(attn.float(), dim=-1).type_as(attn)
|
||||
x = torch.einsum('bnij,bjnc->binc', attn, v)
|
||||
x = x.reshape(b, -1, n * c)
|
||||
return x
|
||||
|
||||
|
||||
class T5Attention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config: T5Config,
|
||||
attn_type: str,
|
||||
attn_type: AttentionType,
|
||||
has_relative_attention_bias=False,
|
||||
cache_config: Optional[CacheConfig] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
@@ -170,6 +179,9 @@ class T5Attention(nn.Module):
|
||||
config.relative_attention_max_distance
|
||||
self.d_model = config.d_model
|
||||
self.key_value_proj_dim = config.d_kv
|
||||
assert cache_config
|
||||
# Alternatively we can get it from kv_cache size in fwd.
|
||||
self.block_size = cache_config.block_size
|
||||
|
||||
# Partition heads across multiple tensor parallel GPUs.
|
||||
tp_world_size = get_tensor_model_parallel_world_size()
|
||||
@@ -187,16 +199,26 @@ class T5Attention(nn.Module):
|
||||
bias=False,
|
||||
quant_config=quant_config)
|
||||
|
||||
self.attn = T5MultiHeadAttention()
|
||||
# NOTE (NickLucche) T5 employs a scaled weight initialization scheme
|
||||
# instead of scaling attention scores directly.
|
||||
self.attn = Attention(self.n_heads,
|
||||
config.d_kv,
|
||||
1.0,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.attn",
|
||||
attn_type=self.attn_type)
|
||||
|
||||
# Only the first SelfAttention block in encoder decoder has this
|
||||
# embedding layer, the others reuse its output.
|
||||
if self.has_relative_attention_bias:
|
||||
self.relative_attention_bias = \
|
||||
VocabParallelEmbedding(self.relative_attention_num_buckets,
|
||||
self.n_heads,
|
||||
org_num_embeddings=self.relative_attention_num_buckets,
|
||||
padding_size=self.relative_attention_num_buckets,
|
||||
org_num_embeddings=\
|
||||
self.relative_attention_num_buckets,
|
||||
quant_config=quant_config)
|
||||
self.o = RowParallelLinear(
|
||||
self.out_proj = RowParallelLinear(
|
||||
self.inner_dim,
|
||||
self.d_model,
|
||||
bias=False,
|
||||
@@ -207,7 +229,7 @@ class T5Attention(nn.Module):
|
||||
def _relative_position_bucket(relative_position,
|
||||
bidirectional=True,
|
||||
num_buckets=32,
|
||||
max_distance=128) -> torch.Tensor:
|
||||
max_distance=128):
|
||||
"""
|
||||
Adapted from Mesh Tensorflow:
|
||||
https://github.com/tensorflow/mesh/blob/0cb87fe07da627bf0b7e60475d59f95ed6b5be3d/mesh_tensorflow/transformer/transformer_layers.py#L593
|
||||
@@ -264,6 +286,7 @@ class T5Attention(nn.Module):
|
||||
key_length,
|
||||
device=None) -> torch.Tensor:
|
||||
"""Compute binned relative position bias"""
|
||||
# TODO possible tp issue?
|
||||
if device is None:
|
||||
device = self.relative_attention_bias.weight.device
|
||||
context_position = torch.arange(query_length,
|
||||
@@ -290,43 +313,110 @@ class T5Attention(nn.Module):
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor, # (num_tokens, d_model)
|
||||
attention_mask: torch.Tensor,
|
||||
attn_metadata: Optional[AttentionMetadata] = None,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
bs, seq_len, _ = hidden_states.shape
|
||||
num_seqs = bs
|
||||
n, c = self.n_heads, self.d_model // self.n_heads
|
||||
# TODO auto-selection of xformers backend when t5 is detected
|
||||
assert isinstance(attn_metadata, XFormersMetadata)
|
||||
num_seqs = len(
|
||||
attn_metadata.seq_lens) if attn_metadata.seq_lens else len(
|
||||
attn_metadata.encoder_seq_lens)
|
||||
qkv, _ = self.qkv_proj(hidden_states)
|
||||
# Projection of 'own' hidden state (self-attention). No GQA here.
|
||||
q, k, v = qkv.split(self.inner_dim, dim=-1)
|
||||
q = q.view(bs, -1, n, c)
|
||||
k = k.view(bs, -1, n, c)
|
||||
v = v.view(bs, -1, n, c)
|
||||
|
||||
assert attn_metadata is not None
|
||||
attn_bias = attn_metadata.attn_bias
|
||||
# NOTE (NickLucche) Attn bias is computed once per encoder or decoder
|
||||
# forward, on the first call to T5Attention.forward. Subsequent
|
||||
# *self-attention* layers will reuse it.
|
||||
attn_bias = _get_attn_bias(attn_metadata, self.attn_type)
|
||||
if self.attn_type == AttentionType.ENCODER_DECODER:
|
||||
# Projection of encoder's hidden states, cross-attention.
|
||||
if encoder_hidden_states is None:
|
||||
# Decode phase, kv already cached
|
||||
assert attn_metadata.num_prefills == 0
|
||||
k = None
|
||||
v = None
|
||||
else:
|
||||
assert attn_metadata.num_prefills > 0
|
||||
# Prefill phase (first decoder forward), caching kv
|
||||
qkv_enc, _ = self.qkv_proj(encoder_hidden_states)
|
||||
_, k, v = qkv_enc.split(self.inner_dim, dim=-1)
|
||||
# No custom attention bias must be set when running cross attn.
|
||||
assert attn_bias is None
|
||||
|
||||
# Not compatible with CP here (as all encoder-decoder models),
|
||||
# as it assumes homogeneous batch (prefills or decodes).
|
||||
if self.has_relative_attention_bias:
|
||||
elif self.has_relative_attention_bias:
|
||||
assert attn_bias is None # to be recomputed
|
||||
# Self-attention. Compute T5 relative positional encoding.
|
||||
# The bias term is computed on longest sequence in batch. Biases
|
||||
# for shorter sequences are slices of the longest.
|
||||
assert self.attn_type == AttentionType.ENCODER
|
||||
attn_bias = self.compute_bias(seq_len,
|
||||
seq_len).repeat(num_seqs, 1, 1, 1)
|
||||
attn_metadata.attn_bias = attn_bias
|
||||
else:
|
||||
# TODO xformers-specific code.
|
||||
align_to = 8
|
||||
# bias expected shape: (num_seqs, NH, L, L_pad) for prefill,
|
||||
# (num_seqs, NH, 1, L_pad) for decodes.
|
||||
if self.attn_type == AttentionType.ENCODER:
|
||||
# Encoder prefill stage, uses xFormers, hence sequence
|
||||
# padding/alignment to 8 is required.
|
||||
seq_len = attn_metadata.max_encoder_seq_len
|
||||
padded_seq_len = (seq_len + align_to -
|
||||
1) // align_to * align_to
|
||||
# TODO (NickLucche) avoid extra copy on repeat,
|
||||
# provide multiple slices of same memory
|
||||
position_bias = self.compute_bias(seq_len,
|
||||
padded_seq_len).repeat(
|
||||
num_seqs, 1, 1, 1)
|
||||
# xFormers expects a list of biases, one matrix per sequence.
|
||||
# As each sequence gets its own bias, no masking is required.
|
||||
attn_bias = [
|
||||
p[None, :, :sq, :sq] for p, sq in zip(
|
||||
position_bias, attn_metadata.encoder_seq_lens)
|
||||
]
|
||||
elif attn_metadata.prefill_metadata:
|
||||
# Decoder prefill stage, uses xFormers, hence sequence
|
||||
# padding/alignment to 8 is required. First decoder step,
|
||||
# seq_len is usually 1, but one can prepend different start
|
||||
# tokens prior to generation.
|
||||
seq_len = attn_metadata.max_prefill_seq_len
|
||||
# ->align
|
||||
padded_seq_len = (seq_len + align_to -
|
||||
1) // align_to * align_to
|
||||
position_bias = self.compute_bias(seq_len,
|
||||
padded_seq_len).repeat(
|
||||
num_seqs, 1, 1, 1)
|
||||
# Causal mask for prefill.
|
||||
attn_bias = [
|
||||
LowerTriangularMaskWithTensorBias(pb[None, :, :sq, :sq])
|
||||
for pb, sq in zip(position_bias, attn_metadata.seq_lens)
|
||||
]
|
||||
else:
|
||||
# Decoder decoding stage, uses PagedAttention, hence sequence
|
||||
# padding/alignment to `block_size` is required. Expected
|
||||
# number of queries is always 1 (MQA not supported).
|
||||
seq_len = attn_metadata.max_decode_seq_len
|
||||
block_aligned_seq_len = (seq_len + self.block_size - 1
|
||||
) // self.block_size * self.block_size
|
||||
|
||||
# TODO bf16 bias support in PagedAttention.
|
||||
position_bias = self.compute_bias(
|
||||
seq_len, block_aligned_seq_len).float()
|
||||
# Bias for the last query, the one at current decoding step.
|
||||
position_bias = position_bias[:, :, -1:, :].repeat(
|
||||
num_seqs, 1, 1, 1)
|
||||
# No explicit masking required, this is done inside the
|
||||
# paged attention kernel based on the sequence length.
|
||||
attn_bias = [position_bias]
|
||||
|
||||
# NOTE Assign bias term on metadata based on attn type:
|
||||
# ENCODER->`encoder_attn_bias`, DECODER->`attn_bias`.
|
||||
_set_attn_bias(attn_metadata, attn_bias, self.attn_type)
|
||||
elif not self.has_relative_attention_bias:
|
||||
# Encoder/Decoder Self-Attention Layer, attn bias already cached.
|
||||
assert attn_bias is not None
|
||||
|
||||
if attention_mask is not None:
|
||||
attention_mask = attention_mask.view(
|
||||
bs, 1, 1,
|
||||
-1) if attention_mask.ndim == 2 else attention_mask.unsqueeze(1)
|
||||
attn_bias.masked_fill_(attention_mask == 0,
|
||||
torch.finfo(q.dtype).min)
|
||||
attn_output = self.attn(q, k, v, attn_bias)
|
||||
output, _ = self.o(attn_output)
|
||||
attn_output = self.attn(q, k, v, kv_cache, attn_metadata)
|
||||
output, _ = self.out_proj(attn_output)
|
||||
return output
|
||||
|
||||
|
||||
@@ -336,6 +426,7 @@ class T5LayerSelfAttention(nn.Module):
|
||||
self,
|
||||
config,
|
||||
has_relative_attention_bias=False,
|
||||
cache_config: Optional[CacheConfig] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
@@ -345,21 +436,24 @@ class T5LayerSelfAttention(nn.Module):
|
||||
AttentionType.DECODER
|
||||
if "decoder" in prefix else AttentionType.ENCODER,
|
||||
has_relative_attention_bias=has_relative_attention_bias,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.SelfAttention")
|
||||
self.layer_norm = RMSNorm(config.d_model, eps=config.layer_norm_epsilon)
|
||||
self.layer_norm = T5LayerNorm(config.d_model,
|
||||
eps=config.layer_norm_epsilon)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
attn_metadata: Optional[AttentionMetadata] = None,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
) -> torch.Tensor:
|
||||
normed_hidden_states = self.layer_norm.forward_native(hidden_states)
|
||||
normed_hidden_states = self.layer_norm(hidden_states)
|
||||
attention_output = self.SelfAttention(
|
||||
hidden_states=normed_hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
kv_cache=kv_cache,
|
||||
attn_metadata=attn_metadata,
|
||||
encoder_hidden_states=None,
|
||||
)
|
||||
hidden_states = hidden_states + attention_output
|
||||
return hidden_states
|
||||
@@ -369,25 +463,32 @@ class T5LayerCrossAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
config,
|
||||
cache_config: Optional[CacheConfig] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
self.EncDecAttention = T5Attention(config,
|
||||
AttentionType.ENCODER_DECODER,
|
||||
has_relative_attention_bias=False,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.EncDecAttention")
|
||||
self.layer_norm = RMSNorm(config.d_model, eps=config.layer_norm_epsilon)
|
||||
self.layer_norm = T5LayerNorm(config.d_model,
|
||||
eps=config.layer_norm_epsilon)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attn_metadata: Optional[AttentionMetadata] = None,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
normed_hidden_states = self.layer_norm.forward_native(hidden_states)
|
||||
normed_hidden_states = self.layer_norm(hidden_states)
|
||||
attention_output = self.EncDecAttention(
|
||||
hidden_states=normed_hidden_states,
|
||||
kv_cache=kv_cache,
|
||||
attn_metadata=attn_metadata,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
)
|
||||
hidden_states = hidden_states + attention_output
|
||||
return hidden_states
|
||||
@@ -399,44 +500,50 @@ class T5Block(nn.Module):
|
||||
config: T5Config,
|
||||
is_decoder: bool,
|
||||
has_relative_attention_bias=False,
|
||||
cache_config: Optional[CacheConfig] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
self.is_decoder = is_decoder
|
||||
self.layer = nn.ModuleList()
|
||||
self.layer.append(
|
||||
T5LayerSelfAttention(
|
||||
config,
|
||||
has_relative_attention_bias=has_relative_attention_bias,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.self_attn"))
|
||||
self.self_attn = T5LayerSelfAttention(
|
||||
config,
|
||||
has_relative_attention_bias=has_relative_attention_bias,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.self_attn")
|
||||
|
||||
if self.is_decoder:
|
||||
self.layer.append(
|
||||
T5LayerCrossAttention(config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.cross_attn"))
|
||||
self.cross_attn = T5LayerCrossAttention(
|
||||
config,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.cross_attn")
|
||||
|
||||
self.layer.append(T5LayerFF(config, quant_config=quant_config))
|
||||
self.ffn = T5LayerFF(config, quant_config=quant_config)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
attn_metadata: Optional[AttentionMetadata] = None,
|
||||
kv_cache: torch.Tensor,
|
||||
attn_metadata: AttentionMetadata,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
|
||||
hidden_states = self.layer[0](hidden_states=hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
attn_metadata=attn_metadata)
|
||||
hidden_states = self.self_attn(
|
||||
hidden_states=hidden_states,
|
||||
kv_cache=kv_cache,
|
||||
attn_metadata=attn_metadata,
|
||||
)
|
||||
if self.is_decoder:
|
||||
hidden_states = self.layer[1](hidden_states=hidden_states,
|
||||
attn_metadata=attn_metadata)
|
||||
hidden_states = self.cross_attn(
|
||||
hidden_states=hidden_states,
|
||||
kv_cache=kv_cache,
|
||||
attn_metadata=attn_metadata,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
)
|
||||
|
||||
# Apply Feed Forward layer
|
||||
hidden_states = self.layer[2](hidden_states)
|
||||
else:
|
||||
hidden_states = self.layer[1](hidden_states)
|
||||
# Apply Feed Forward layer
|
||||
hidden_states = self.ffn(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
@@ -447,221 +554,234 @@ class T5Stack(nn.Module):
|
||||
is_decoder: bool,
|
||||
n_layers: int,
|
||||
embed_tokens=None,
|
||||
cache_config: Optional[CacheConfig] = None,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
is_umt5: bool = False):
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
self.embed_tokens = embed_tokens
|
||||
self.is_umt5 = is_umt5
|
||||
if is_umt5:
|
||||
self.block = nn.ModuleList([
|
||||
T5Block(config,
|
||||
is_decoder=is_decoder,
|
||||
has_relative_attention_bias=True,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.blocks.{i}") for i in range(n_layers)
|
||||
])
|
||||
# Only the first block has relative positional encoding.
|
||||
self.blocks = nn.ModuleList([
|
||||
T5Block(config,
|
||||
is_decoder=is_decoder,
|
||||
has_relative_attention_bias=i == 0,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.blocks.{i}") for i in range(n_layers)
|
||||
])
|
||||
self.final_layer_norm = T5LayerNorm(config.d_model,
|
||||
eps=config.layer_norm_epsilon)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
kv_caches: List[torch.Tensor],
|
||||
attn_metadata: AttentionMetadata,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None
|
||||
) -> torch.Tensor:
|
||||
hidden_states = self.embed_tokens(input_ids)
|
||||
|
||||
for idx, block in enumerate(self.blocks):
|
||||
hidden_states = block(
|
||||
hidden_states=hidden_states,
|
||||
kv_cache=kv_caches[idx],
|
||||
attn_metadata=attn_metadata,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
)
|
||||
hidden_states = self.final_layer_norm(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class T5Model(nn.Module):
|
||||
_tied_weights_keys = [
|
||||
"encoder.embed_tokens.weight", "decoder.embed_tokens.weight"
|
||||
]
|
||||
|
||||
def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
|
||||
super().__init__()
|
||||
config: T5Config = vllm_config.model_config.hf_config
|
||||
cache_config = vllm_config.cache_config
|
||||
quant_config = vllm_config.quant_config
|
||||
lora_config = vllm_config.lora_config
|
||||
|
||||
lora_vocab = (lora_config.lora_extra_vocab_size *
|
||||
(lora_config.max_loras or 1)) if lora_config else 0
|
||||
self.vocab_size = config.vocab_size + lora_vocab
|
||||
self.padding_idx = config.pad_token_id
|
||||
self.shared = VocabParallelEmbedding(
|
||||
config.vocab_size,
|
||||
config.d_model,
|
||||
org_num_embeddings=config.vocab_size)
|
||||
|
||||
self.encoder = T5Stack(config,
|
||||
False,
|
||||
config.num_layers,
|
||||
self.shared,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.encoder")
|
||||
self.decoder = T5Stack(config,
|
||||
True,
|
||||
config.num_decoder_layers,
|
||||
self.shared,
|
||||
cache_config=cache_config,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.decoder")
|
||||
|
||||
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
|
||||
return self.shared(input_ids)
|
||||
|
||||
def forward(self, input_ids: torch.Tensor, encoder_input_ids: torch.Tensor,
|
||||
kv_caches: List[torch.Tensor],
|
||||
attn_metadata: AttentionMetadata) -> torch.Tensor:
|
||||
encoder_hidden_states = None
|
||||
|
||||
if encoder_input_ids.numel() > 0:
|
||||
# Run encoder attention if a non-zero number of encoder tokens
|
||||
# are provided as input: on a regular generate call, the encoder
|
||||
# runs once, on the prompt. Subsequent decoder calls reuse output
|
||||
# `encoder_hidden_states`.
|
||||
encoder_hidden_states = self.encoder(input_ids=encoder_input_ids,
|
||||
kv_caches=kv_caches,
|
||||
attn_metadata=attn_metadata)
|
||||
# Clear attention bias state.
|
||||
attn_metadata.attn_bias = None
|
||||
attn_metadata.encoder_attn_bias = None
|
||||
attn_metadata.cross_attn_bias = None
|
||||
|
||||
decoder_outputs = self.decoder(
|
||||
input_ids=input_ids,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
kv_caches=kv_caches,
|
||||
attn_metadata=attn_metadata)
|
||||
|
||||
# When capturing CUDA Graph
|
||||
attn_metadata.attn_bias = None
|
||||
attn_metadata.encoder_attn_bias = None
|
||||
attn_metadata.cross_attn_bias = None
|
||||
return decoder_outputs
|
||||
|
||||
|
||||
class T5ForConditionalGeneration(nn.Module):
|
||||
_keys_to_ignore_on_load_unexpected = [
|
||||
"decoder.block.0.layer.1.EncDecAttention.relative_attention_bias.weight",
|
||||
]
|
||||
_tied_weights_keys = [
|
||||
"encoder.embed_tokens.weight", "decoder.embed_tokens.weight",
|
||||
"lm_head.weight"
|
||||
]
|
||||
|
||||
def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""):
|
||||
super().__init__()
|
||||
config: T5Config = vllm_config.model_config.hf_config
|
||||
self.model_dim = config.d_model
|
||||
self.config = config
|
||||
self.unpadded_vocab_size = config.vocab_size
|
||||
if lora_config := vllm_config.lora_config:
|
||||
self.unpadded_vocab_size += lora_config.lora_extra_vocab_size
|
||||
|
||||
self.model = T5Model(vllm_config=vllm_config,
|
||||
prefix=maybe_prefix(prefix, "model"))
|
||||
# Although not in config, this is the default for hf models.
|
||||
if self.config.tie_word_embeddings:
|
||||
self.lm_head = self.model.shared
|
||||
# in transformers this is smt more explicit, as in (after load)
|
||||
# self.lm_head.weight = self.model.shared.weight
|
||||
else:
|
||||
# Only the first block has relative positional encoding.
|
||||
self.block = nn.ModuleList([
|
||||
T5Block(config,
|
||||
is_decoder=is_decoder,
|
||||
has_relative_attention_bias=i == 0,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.blocks.{i}") for i in range(n_layers)
|
||||
])
|
||||
self.final_layer_norm = RMSNorm(config.d_model,
|
||||
eps=config.layer_norm_epsilon)
|
||||
self.lm_head = ParallelLMHead(self.unpadded_vocab_size,
|
||||
config.d_model,
|
||||
org_num_embeddings=config.vocab_size)
|
||||
|
||||
self.logits_processor = LogitsProcessor(self.unpadded_vocab_size,
|
||||
config.vocab_size)
|
||||
self.sampler = get_sampler()
|
||||
|
||||
def compute_logits(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
sampling_metadata: SamplingMetadata,
|
||||
) -> Optional[torch.Tensor]:
|
||||
if self.config.tie_word_embeddings:
|
||||
# Rescale output before projecting on vocab
|
||||
# See https://github.com/tensorflow/mesh/blob/fa19d69eafc9a482aff0b59ddd96b025c0cb207d/mesh_tensorflow/transformer/transformer.py#L586 # noqa: E501
|
||||
hidden_states = hidden_states * (self.model_dim**-0.5)
|
||||
logits = self.logits_processor(self.lm_head, hidden_states,
|
||||
sampling_metadata)
|
||||
return logits
|
||||
|
||||
def sample(
|
||||
self,
|
||||
logits: Optional[torch.Tensor],
|
||||
sampling_metadata: SamplingMetadata,
|
||||
) -> Optional[SamplerOutput]:
|
||||
next_tokens = self.sampler(logits, sampling_metadata)
|
||||
return next_tokens
|
||||
|
||||
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
|
||||
return self.model.shared(input_ids)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
positions: torch.Tensor,
|
||||
kv_caches: List[torch.Tensor],
|
||||
attn_metadata: AttentionMetadata,
|
||||
intermediate_tensors: Optional[IntermediateTensors] = None,
|
||||
*,
|
||||
encoder_input_ids: torch.Tensor,
|
||||
encoder_positions: torch.Tensor,
|
||||
**kwargs,
|
||||
) -> torch.Tensor:
|
||||
hidden_states = self.embed_tokens(input_ids)
|
||||
return self.model(input_ids, encoder_input_ids, kv_caches,
|
||||
attn_metadata)
|
||||
|
||||
for idx, block in enumerate(self.block):
|
||||
hidden_states = block(
|
||||
hidden_states=hidden_states,
|
||||
attention_mask=attention_mask,
|
||||
attn_metadata=attn_metadata,
|
||||
)
|
||||
hidden_states = self.final_layer_norm.forward_native(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class T5EncoderModel(nn.Module):
|
||||
|
||||
def __init__(self, config: T5Config, prefix: str = ""):
|
||||
super().__init__()
|
||||
|
||||
quant_config = None
|
||||
|
||||
self.shared = VocabParallelEmbedding(
|
||||
config.vocab_size,
|
||||
config.d_model,
|
||||
org_num_embeddings=config.vocab_size)
|
||||
|
||||
self.encoder = T5Stack(config,
|
||||
False,
|
||||
config.num_layers,
|
||||
self.shared,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.encoder",
|
||||
is_umt5=False)
|
||||
|
||||
def get_input_embeddings(self):
|
||||
return self.shared
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.LongTensor] = None,
|
||||
attention_mask: Optional[torch.FloatTensor] = None,
|
||||
head_mask: Optional[torch.FloatTensor] = None,
|
||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||
output_attentions: Optional[bool] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
return_dict: Optional[bool] = None,
|
||||
) -> torch.Tensor:
|
||||
attn_metadata = AttentionMetadata(None)
|
||||
encoder_outputs = self.encoder(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
attn_metadata=attn_metadata,
|
||||
)
|
||||
|
||||
return encoder_outputs
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]):
|
||||
model_params_dict = dict(self.named_parameters(remove_duplicate=False))
|
||||
loaded_params: Set[str] = set()
|
||||
renamed_reg = [
|
||||
(re.compile(r'block\.(\d+)\.layer\.0'), r'blocks.\1.self_attn'),
|
||||
(re.compile(r'decoder.block\.(\d+)\.layer\.1'),
|
||||
r'decoder.blocks.\1.cross_attn'),
|
||||
(re.compile(r'decoder.block\.(\d+)\.layer\.2'),
|
||||
r'decoder.blocks.\1.ffn'),
|
||||
# encoder has no cross-attn, but rather self-attention+ffn.
|
||||
(re.compile(r'encoder.block\.(\d+)\.layer\.1'),
|
||||
r'encoder.blocks.\1.ffn'),
|
||||
(re.compile(r'\.o\.'), r'.out_proj.'),
|
||||
]
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q", "q"),
|
||||
(".qkv_proj", ".k", "k"),
|
||||
(".qkv_proj", ".v", "v"),
|
||||
(".qkv_proj.", ".q.", "q"),
|
||||
(".qkv_proj.", ".k.", "k"),
|
||||
(".qkv_proj.", ".v.", "v")
|
||||
]
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: Set[str] = set()
|
||||
|
||||
for name, loaded_weight in weights:
|
||||
loaded = False
|
||||
if "decoder" in name or "lm_head" in name:
|
||||
# No relative position attn bias on cross attention.
|
||||
if name in self._keys_to_ignore_on_load_unexpected:
|
||||
continue
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
|
||||
# Handle some renaming
|
||||
for reg in renamed_reg:
|
||||
name = re.sub(*reg, name)
|
||||
|
||||
top_module, _ = name.split('.', 1)
|
||||
if top_module != 'lm_head':
|
||||
name = f"model.{name}"
|
||||
|
||||
# Split q/k/v layers to unified QKVParallelLinear
|
||||
for (param_name, weight_name, shard_id) in stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
|
||||
if name not in params_dict:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
param = model_params_dict[name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
loaded = True
|
||||
break
|
||||
if not loaded:
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
|
||||
if name not in params_dict:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
else:
|
||||
# Not a q/k/v layer.
|
||||
param = model_params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader",
|
||||
default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(name)
|
||||
return loaded_params
|
||||
|
||||
|
||||
class UMT5EncoderModel(nn.Module):
|
||||
|
||||
def __init__(self, config: T5Config, prefix: str = ""):
|
||||
super().__init__()
|
||||
|
||||
quant_config = None
|
||||
|
||||
self.shared = VocabParallelEmbedding(
|
||||
config.vocab_size,
|
||||
config.d_model,
|
||||
org_num_embeddings=config.vocab_size)
|
||||
|
||||
self.encoder = T5Stack(config,
|
||||
False,
|
||||
config.num_layers,
|
||||
self.shared,
|
||||
quant_config=quant_config,
|
||||
prefix=f"{prefix}.encoder",
|
||||
is_umt5=True)
|
||||
|
||||
def get_input_embeddings(self):
|
||||
return self.shared
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: Optional[torch.LongTensor] = None,
|
||||
attention_mask: Optional[torch.FloatTensor] = None,
|
||||
head_mask: Optional[torch.FloatTensor] = None,
|
||||
inputs_embeds: Optional[torch.FloatTensor] = None,
|
||||
output_attentions: Optional[bool] = None,
|
||||
output_hidden_states: Optional[bool] = None,
|
||||
return_dict: Optional[bool] = None,
|
||||
) -> torch.Tensor:
|
||||
attn_metadata = AttentionMetadata(None)
|
||||
encoder_outputs = self.encoder(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
attn_metadata=attn_metadata,
|
||||
)
|
||||
|
||||
return encoder_outputs
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q", "q"),
|
||||
(".qkv_proj", ".k", "k"),
|
||||
(".qkv_proj", ".v", "v"),
|
||||
]
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: Set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
loaded = False
|
||||
if "decoder" in name or "lm_head" in name:
|
||||
continue
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
|
||||
if name not in params_dict:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
weight_loader = param.weight_loader
|
||||
weight_loader(param, loaded_weight, shard_id)
|
||||
loaded = True
|
||||
break
|
||||
if not loaded:
|
||||
# Skip loading extra bias for GPTQ models.
|
||||
if name.endswith(".bias") and name not in params_dict:
|
||||
continue
|
||||
|
||||
if name not in params_dict:
|
||||
continue
|
||||
|
||||
param = params_dict[name]
|
||||
weight_loader = getattr(param, "weight_loader",
|
||||
default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add(name)
|
||||
return loaded_params
|
||||
return loaded_params
|
||||
@@ -0,0 +1,22 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from typing import List
|
||||
|
||||
def extract_layer_index(layer_name: str) -> int:
|
||||
"""
|
||||
Extract the layer index from the module name.
|
||||
Examples:
|
||||
- "encoder.layers.0" -> 0
|
||||
- "encoder.layers.1.self_attn" -> 1
|
||||
- "2.self_attn" -> 2
|
||||
- "model.encoder.layers.0.sub.1" -> ValueError
|
||||
"""
|
||||
subnames = layer_name.split(".")
|
||||
int_vals: List[int] = []
|
||||
for subname in subnames:
|
||||
try:
|
||||
int_vals.append(int(subname))
|
||||
except ValueError:
|
||||
continue
|
||||
assert len(int_vals) == 1, (f"layer name {layer_name} should"
|
||||
" only contain one integer")
|
||||
return int_vals[0]
|
||||
@@ -1,13 +1,16 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/models/vision.py
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Generic, Optional, TypeVar, Union
|
||||
from typing import Final, Generic, Optional, Protocol, TypeVar, Union
|
||||
|
||||
import torch
|
||||
from transformers import PretrainedConfig
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from vllm.attention.selector import (backend_name_to_enum,
|
||||
get_global_forced_attn_backend)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.platforms import _Backend, current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -47,6 +50,62 @@ class VisionEncoderInfo(ABC, Generic[_C]):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class VisionLanguageConfig(Protocol):
|
||||
vision_config: Final[PretrainedConfig]
|
||||
|
||||
|
||||
def get_vision_encoder_info(
|
||||
hf_config: VisionLanguageConfig) -> VisionEncoderInfo:
|
||||
# Avoid circular imports
|
||||
from .clip import CLIPEncoderInfo, CLIPVisionConfig
|
||||
from .pixtral import PixtralHFEncoderInfo, PixtralVisionConfig
|
||||
from .siglip import SiglipEncoderInfo, SiglipVisionConfig
|
||||
|
||||
vision_config = hf_config.vision_config
|
||||
if isinstance(vision_config, CLIPVisionConfig):
|
||||
return CLIPEncoderInfo(vision_config)
|
||||
if isinstance(vision_config, PixtralVisionConfig):
|
||||
return PixtralHFEncoderInfo(vision_config)
|
||||
if isinstance(vision_config, SiglipVisionConfig):
|
||||
return SiglipEncoderInfo(vision_config)
|
||||
|
||||
msg = f"Unsupported vision config: {type(vision_config)}"
|
||||
raise NotImplementedError(msg)
|
||||
|
||||
|
||||
def get_vit_attn_backend(support_fa: bool = False) -> _Backend:
|
||||
"""
|
||||
Get the available attention backend for Vision Transformer.
|
||||
"""
|
||||
# TODO(Isotr0py): Remove `support_fa` after support FA for all ViTs attn.
|
||||
selected_backend: Optional[_Backend] = get_global_forced_attn_backend()
|
||||
if selected_backend is None:
|
||||
backend_by_env_var: Optional[str] = envs.VLLM_ATTENTION_BACKEND
|
||||
if backend_by_env_var is not None:
|
||||
selected_backend = backend_name_to_enum(backend_by_env_var)
|
||||
if selected_backend is None:
|
||||
if current_platform.is_cuda():
|
||||
device_available = current_platform.has_device_capability(80)
|
||||
if device_available and support_fa:
|
||||
from transformers.utils import is_flash_attn_2_available
|
||||
if is_flash_attn_2_available():
|
||||
selected_backend = _Backend.FLASH_ATTN
|
||||
else:
|
||||
logger.warning_once(
|
||||
"Current `vllm-flash-attn` has a bug inside vision "
|
||||
"module, so we use xformers backend instead. You can "
|
||||
"run `pip install flash-attn` to use flash-attention "
|
||||
"backend.")
|
||||
selected_backend = _Backend.XFORMERS
|
||||
else:
|
||||
# For Volta and Turing GPUs, use xformers instead.
|
||||
selected_backend = _Backend.XFORMERS
|
||||
else:
|
||||
# Default to torch SDPA for other non-GPU platforms.
|
||||
selected_backend = _Backend.TORCH_SDPA
|
||||
return selected_backend
|
||||
|
||||
|
||||
def resolve_visual_encoder_outputs(
|
||||
encoder_outputs: Union[torch.Tensor, list[torch.Tensor]],
|
||||
feature_sample_layers: Optional[list[int]],
|
||||
|
||||
@@ -1,6 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from SGLang: https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/hf_transformers_utils.py
|
||||
|
||||
# Copyright 2023-2024 SGLang Team
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
@@ -17,15 +14,24 @@
|
||||
"""Utilities for Huggingface Transformers."""
|
||||
|
||||
import contextlib
|
||||
import json
|
||||
import os
|
||||
import warnings
|
||||
from pathlib import Path
|
||||
from typing import Any, Dict, Optional, Type, Union
|
||||
from typing import Dict, Optional, Type, Union, Any
|
||||
import json
|
||||
|
||||
from huggingface_hub import snapshot_download
|
||||
from transformers import AutoConfig, PretrainedConfig
|
||||
from transformers.models.auto.modeling_auto import (
|
||||
MODEL_FOR_CAUSAL_LM_MAPPING_NAMES)
|
||||
from transformers import (
|
||||
AutoConfig,
|
||||
AutoProcessor,
|
||||
AutoTokenizer,
|
||||
PretrainedConfig,
|
||||
PreTrainedTokenizer,
|
||||
PreTrainedTokenizerFast,
|
||||
)
|
||||
from transformers.models.auto.modeling_auto import MODEL_FOR_CAUSAL_LM_MAPPING_NAMES
|
||||
|
||||
# from fastvideo.v1.models.configs import ChatGLMConfig, DbrxConfig, ExaoneConfig, Qwen2_5_VLConfig
|
||||
|
||||
_CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = {
|
||||
# ChatGLMConfig.model_type: ChatGLMConfig,
|
||||
@@ -43,8 +49,7 @@ def download_from_hf(model_path: str):
|
||||
if os.path.exists(model_path):
|
||||
return model_path
|
||||
|
||||
return snapshot_download(model_path,
|
||||
allow_patterns=["*.json", "*.bin", "*.model"])
|
||||
return snapshot_download(model_path, allow_patterns=["*.json", "*.bin", "*.model"])
|
||||
|
||||
|
||||
def get_hf_config(
|
||||
@@ -57,25 +62,24 @@ def get_hf_config(
|
||||
):
|
||||
is_gguf = check_gguf_file(model)
|
||||
if is_gguf:
|
||||
raise NotImplementedError("GGUF models are not supported.")
|
||||
kwargs["gguf_file"] = model
|
||||
model = Path(model).parent
|
||||
|
||||
config = AutoConfig.from_pretrained(model,
|
||||
trust_remote_code=trust_remote_code,
|
||||
revision=revision,
|
||||
**kwargs)
|
||||
config = AutoConfig.from_pretrained(
|
||||
model, trust_remote_code=trust_remote_code, revision=revision, **kwargs
|
||||
)
|
||||
if config.model_type in _CONFIG_REGISTRY:
|
||||
config_class = _CONFIG_REGISTRY[config.model_type]
|
||||
config = config_class.from_pretrained(model, revision=revision)
|
||||
# NOTE(HandH1998): Qwen2VL requires `_name_or_path` attribute in `config`.
|
||||
config._name_or_path = model
|
||||
setattr(config, "_name_or_path", model)
|
||||
if model_override_args:
|
||||
config.update(model_override_args)
|
||||
|
||||
# Special architecture mapping check for GGUF models
|
||||
if is_gguf:
|
||||
if config.model_type not in MODEL_FOR_CAUSAL_LM_MAPPING_NAMES:
|
||||
raise RuntimeError(
|
||||
f"Can't get gguf config for {config.model_type}.")
|
||||
raise RuntimeError(f"Can't get gguf config for {config.model_type}.")
|
||||
model_type = MODEL_FOR_CAUSAL_LM_MAPPING_NAMES[config.model_type]
|
||||
config.update({"architectures": [model_type]})
|
||||
|
||||
@@ -101,16 +105,13 @@ def get_diffusers_config(
|
||||
if os.path.exists(config_file):
|
||||
try:
|
||||
# Load the config directly from the file
|
||||
with open(config_file) as f:
|
||||
config_dict: Dict[str, Any] = json.load(f)
|
||||
|
||||
with open(config_file, "r") as f:
|
||||
config_dict = json.load(f)
|
||||
|
||||
# TODO(will): apply any overrides from inference args
|
||||
return config_dict
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
f"Failed to load diffusers config from {config_file}: {e}"
|
||||
) from e
|
||||
raise RuntimeError(f"Config file not found at {config_file}")
|
||||
raise RuntimeError(f"Failed to load diffusers config from {config_file}: {e}")
|
||||
else:
|
||||
raise RuntimeError(f"Diffusers config file not found at {model}")
|
||||
|
||||
@@ -128,11 +129,119 @@ CONTEXT_LENGTH_KEYS = [
|
||||
]
|
||||
|
||||
|
||||
def get_context_length(config):
|
||||
"""Get the context length of a model from a huggingface model configs."""
|
||||
text_config = config
|
||||
rope_scaling = getattr(text_config, "rope_scaling", None)
|
||||
if rope_scaling:
|
||||
rope_scaling_factor = rope_scaling.get("factor", 1)
|
||||
if "original_max_position_embeddings" in rope_scaling:
|
||||
rope_scaling_factor = 1
|
||||
if rope_scaling.get("rope_type", None) == "llama3":
|
||||
rope_scaling_factor = 1
|
||||
else:
|
||||
rope_scaling_factor = 1
|
||||
|
||||
for key in CONTEXT_LENGTH_KEYS:
|
||||
val = getattr(text_config, key, None)
|
||||
if val is not None:
|
||||
return int(rope_scaling_factor * val)
|
||||
return 2048
|
||||
|
||||
|
||||
# A fast LLaMA tokenizer with the pre-processed `tokenizer.json` file.
|
||||
_FAST_LLAMA_TOKENIZER = "hf-internal-testing/llama-tokenizer"
|
||||
|
||||
|
||||
def get_tokenizer(
|
||||
tokenizer_name: str,
|
||||
*args,
|
||||
tokenizer_mode: str = "auto",
|
||||
trust_remote_code: bool = False,
|
||||
tokenizer_revision: Optional[str] = None,
|
||||
**kwargs,
|
||||
) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast]:
|
||||
"""Gets a tokenizer for the given model name via Huggingface."""
|
||||
if tokenizer_mode == "slow":
|
||||
if kwargs.get("use_fast", False):
|
||||
raise ValueError("Cannot use the fast tokenizer in slow tokenizer mode.")
|
||||
kwargs["use_fast"] = False
|
||||
|
||||
is_gguf = check_gguf_file(tokenizer_name)
|
||||
if is_gguf:
|
||||
kwargs["gguf_file"] = tokenizer_name
|
||||
tokenizer_name = Path(tokenizer_name).parent
|
||||
|
||||
try:
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
tokenizer_name,
|
||||
*args,
|
||||
trust_remote_code=trust_remote_code,
|
||||
tokenizer_revision=tokenizer_revision,
|
||||
clean_up_tokenization_spaces=False,
|
||||
**kwargs,
|
||||
)
|
||||
except TypeError as e:
|
||||
# The LLaMA tokenizer causes a protobuf error in some environments.
|
||||
err_msg = (
|
||||
"Failed to load the tokenizer. If you are using a LLaMA V1 model "
|
||||
f"consider using '{_FAST_LLAMA_TOKENIZER}' instead of the "
|
||||
"original tokenizer."
|
||||
)
|
||||
raise RuntimeError(err_msg) from e
|
||||
except ValueError as e:
|
||||
# If the error pertains to the tokenizer class not existing or not
|
||||
# currently being imported, suggest using the --trust-remote-code flag.
|
||||
if not trust_remote_code and (
|
||||
"does not exist or is not currently imported." in str(e)
|
||||
or "requires you to execute the tokenizer file" in str(e)
|
||||
):
|
||||
err_msg = (
|
||||
"Failed to load the tokenizer. If the tokenizer is a custom "
|
||||
"tokenizer not yet available in the HuggingFace transformers "
|
||||
"library, consider setting `trust_remote_code=True` in LLM "
|
||||
"or using the `--trust-remote-code` flag in the CLI."
|
||||
)
|
||||
raise RuntimeError(err_msg) from e
|
||||
else:
|
||||
raise e
|
||||
|
||||
if not isinstance(tokenizer, PreTrainedTokenizerFast):
|
||||
warnings.warn(
|
||||
"Using a slow tokenizer. This might cause a significant "
|
||||
"slowdown. Consider using a fast tokenizer instead."
|
||||
)
|
||||
|
||||
attach_additional_stop_token_ids(tokenizer)
|
||||
return tokenizer
|
||||
|
||||
|
||||
def get_processor(
|
||||
tokenizer_name: str,
|
||||
*args,
|
||||
tokenizer_mode: str = "auto",
|
||||
trust_remote_code: bool = False,
|
||||
tokenizer_revision: Optional[str] = None,
|
||||
**kwargs,
|
||||
):
|
||||
processor = AutoProcessor.from_pretrained(
|
||||
tokenizer_name,
|
||||
*args,
|
||||
trust_remote_code=trust_remote_code,
|
||||
tokenizer_revision=tokenizer_revision,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
attach_additional_stop_token_ids(processor.tokenizer)
|
||||
return processor
|
||||
|
||||
|
||||
def attach_additional_stop_token_ids(tokenizer):
|
||||
# Special handling for stop token <|eom_id|> generated by llama 3 tool use.
|
||||
if "<|eom_id|>" in tokenizer.get_added_vocab():
|
||||
tokenizer.additional_stop_token_ids = set(
|
||||
[tokenizer.get_added_vocab()["<|eom_id|>"]])
|
||||
[tokenizer.get_added_vocab()["<|eom_id|>"]]
|
||||
)
|
||||
else:
|
||||
tokenizer.additional_stop_token_ids = None
|
||||
|
||||
|
||||
@@ -1,41 +1,40 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import dataclasses
|
||||
import glob
|
||||
import os
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Generator, Iterable, List, Optional, Tuple, cast
|
||||
import dataclasses
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
from transformers import AutoTokenizer, PretrainedConfig
|
||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
|
||||
from fastvideo.v1.inference_args import InferenceArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.hf_transformer_utils import (get_diffusers_config,
|
||||
get_hf_config)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
import os
|
||||
import glob
|
||||
from fastvideo.v1.models.loader.fsdp_load import load_fsdp_model
|
||||
from fastvideo.v1.models.loader.utils import set_default_torch_dtype
|
||||
from transformers import PretrainedConfig, AutoTokenizer
|
||||
from fastvideo.v1.models.hf_transformer_utils import get_hf_config, get_diffusers_config
|
||||
from fastvideo.v1.models import get_scheduler
|
||||
from fastvideo.v1.models.registry import ModelRegistry
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
from typing import Tuple, List, Optional, Any, Generator
|
||||
import time
|
||||
import torch.nn as nn
|
||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
from fastvideo.v1.models.loader.weight_utils import (
|
||||
filter_duplicate_safetensors_files, filter_files_not_needed_for_inference,
|
||||
pt_weights_iterator, safetensors_weights_iterator)
|
||||
from fastvideo.v1.models.registry import ModelRegistry
|
||||
pt_weights_iterator,
|
||||
safetensors_weights_iterator)
|
||||
from fastvideo.v1.models.loader.utils import set_default_torch_dtype
|
||||
from typing import (Any, Dict, Generator, Iterable, List, Optional,
|
||||
Tuple, cast)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ComponentLoader(ABC):
|
||||
"""Base class for loading a specific type of model component."""
|
||||
|
||||
def __init__(self, device=None) -> None:
|
||||
|
||||
def __init__(self, device=None):
|
||||
self.device = device
|
||||
|
||||
|
||||
@abstractmethod
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
|
||||
"""
|
||||
Load the component based on the model path, architecture, and inference args.
|
||||
|
||||
@@ -48,10 +47,9 @@ class ComponentLoader(ABC):
|
||||
The loaded component
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@classmethod
|
||||
def for_module_type(cls, module_type: str,
|
||||
transformers_or_diffusers: str) -> 'ComponentLoader':
|
||||
def for_module_type(cls, module_type: str, transformers_or_diffusers: str) -> 'ComponentLoader':
|
||||
"""
|
||||
Factory method to create a component loader for a specific module type.
|
||||
|
||||
@@ -72,23 +70,20 @@ class ComponentLoader(ABC):
|
||||
"tokenizer": (TokenizerLoader, "transformers"),
|
||||
"tokenizer_2": (TokenizerLoader, "transformers"),
|
||||
}
|
||||
|
||||
|
||||
if module_type in module_loaders:
|
||||
loader_cls, expected_library = module_loaders[module_type]
|
||||
# Assert that the library matches what's expected for this module type
|
||||
assert transformers_or_diffusers == expected_library, f"{module_type} must be loaded from {expected_library}, got {transformers_or_diffusers}"
|
||||
return loader_cls()
|
||||
|
||||
|
||||
# For unknown module types, use a generic loader
|
||||
logger.warning(
|
||||
"No specific loader found for module type: %s. Using generic loader.",
|
||||
module_type)
|
||||
logger.warning(f"No specific loader found for module type: {module_type}. Using generic loader.")
|
||||
return GenericComponentLoader(transformers_or_diffusers)
|
||||
|
||||
|
||||
|
||||
class TextEncoderLoader(ComponentLoader):
|
||||
"""Loader for text encoders."""
|
||||
|
||||
@dataclasses.dataclass
|
||||
class Source:
|
||||
"""A source for weights."""
|
||||
@@ -117,8 +112,8 @@ class TextEncoderLoader(ComponentLoader):
|
||||
"""Prepare weights for the model.
|
||||
|
||||
If the model is not local, it will be downloaded."""
|
||||
# model_name_or_path = (self._maybe_download_from_modelscope(
|
||||
# model_name_or_path, revision) or model_name_or_path)
|
||||
# model_name_or_path = (self._maybe_download_from_modelscope(
|
||||
# model_name_or_path, revision) or model_name_or_path)
|
||||
|
||||
is_local = os.path.isdir(model_name_or_path)
|
||||
assert is_local, "Model path must be a local directory"
|
||||
@@ -127,12 +122,14 @@ class TextEncoderLoader(ComponentLoader):
|
||||
index_file = SAFE_WEIGHTS_INDEX_NAME
|
||||
allow_patterns = ["*.safetensors", "*.bin"]
|
||||
|
||||
|
||||
if fall_back_to_pt:
|
||||
allow_patterns += ["*.pt"]
|
||||
|
||||
if allow_patterns_overrides is not None:
|
||||
allow_patterns = allow_patterns_overrides
|
||||
|
||||
|
||||
hf_folder = model_name_or_path
|
||||
|
||||
hf_weights_files: List[str] = []
|
||||
@@ -168,6 +165,7 @@ class TextEncoderLoader(ComponentLoader):
|
||||
else:
|
||||
weights_iterator = pt_weights_iterator(hf_weights_files)
|
||||
|
||||
|
||||
if self.counter_before_loading_weights == 0.0:
|
||||
self.counter_before_loading_weights = time.perf_counter()
|
||||
# Apply the prefix.
|
||||
@@ -176,13 +174,14 @@ class TextEncoderLoader(ComponentLoader):
|
||||
|
||||
def _get_all_weights(
|
||||
self,
|
||||
model_config: Any,
|
||||
model_config: Dict[str, Any],
|
||||
model: nn.Module,
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
primary_weights = TextEncoderLoader.Source(
|
||||
model_config.model,
|
||||
prefix="",
|
||||
fall_back_to_pt=getattr(model, "fall_back_to_pt_during_load", True),
|
||||
fall_back_to_pt=getattr(model, "fall_back_to_pt_during_load",
|
||||
True),
|
||||
allow_patterns_overrides=getattr(model, "allow_patterns_overrides",
|
||||
None),
|
||||
)
|
||||
@@ -194,9 +193,8 @@ class TextEncoderLoader(ComponentLoader):
|
||||
)
|
||||
for source in secondary_weights:
|
||||
yield from self._get_weights_iterator(source)
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
|
||||
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
|
||||
"""Load the text encoders based on the model path, architecture, and inference args."""
|
||||
model_config: PretrainedConfig = get_hf_config(
|
||||
model=model_path,
|
||||
@@ -205,20 +203,20 @@ class TextEncoderLoader(ComponentLoader):
|
||||
model_override_args=None,
|
||||
inference_args=inference_args,
|
||||
)
|
||||
logger.info("HF Model config: %s", model_config)
|
||||
|
||||
logger.info(f"HF Model config: {model_config}")
|
||||
|
||||
|
||||
target_device = torch.device(inference_args.device_str)
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(model_path, model_config, target_device)
|
||||
|
||||
def load_model(self, model_path: str, model_config,
|
||||
target_device: torch.device):
|
||||
|
||||
def load_model(self, model_path: str, model_config, target_device: torch.device):
|
||||
with set_default_torch_dtype(torch.float16):
|
||||
with target_device:
|
||||
architectures = getattr(model_config, "architectures", [])
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
|
||||
model = model_cls(model_config)
|
||||
|
||||
|
||||
weights_to_load = {name for name, _ in model.named_parameters()}
|
||||
model_config.model = model_path
|
||||
loaded_weights = model.load_weights(
|
||||
@@ -233,42 +231,41 @@ class TextEncoderLoader(ComponentLoader):
|
||||
# if loaded_weights is not None:
|
||||
weights_not_loaded = weights_to_load - loaded_weights
|
||||
if weights_not_loaded:
|
||||
raise ValueError("Following weights were not initialized from "
|
||||
f"checkpoint: {weights_not_loaded}")
|
||||
raise ValueError(
|
||||
"Following weights were not initialized from "
|
||||
f"checkpoint: {weights_not_loaded}")
|
||||
|
||||
# TODO(will): add support for training/finetune
|
||||
return model.eval()
|
||||
|
||||
|
||||
|
||||
class TokenizerLoader(ComponentLoader):
|
||||
"""Loader for tokenizers."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
|
||||
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
|
||||
"""Load the tokenizer based on the model path, architecture, and inference args."""
|
||||
logger.info("Loading tokenizer from %s", model_path)
|
||||
|
||||
logger.info(f"Loading tokenizer from {model_path}")
|
||||
|
||||
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
model_path,
|
||||
# TODO(will): pass these tokenizer kwargs from inference args? Maybe
|
||||
# other method of config?
|
||||
padding_size='right',
|
||||
)
|
||||
logger.info("Loaded tokenizer: %s", tokenizer.__class__.__name__)
|
||||
logger.info(f"Loaded tokenizer: {tokenizer.__class__.__name__}")
|
||||
return tokenizer
|
||||
|
||||
|
||||
class VAELoader(ComponentLoader):
|
||||
"""Loader for VAE."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
|
||||
"""Load the VAE based on the model path, architecture, and inference args."""
|
||||
# TODO(will): move this to a constants file
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
config = get_diffusers_config(model=model_path)
|
||||
|
||||
|
||||
class_name = config.pop("_class_name")
|
||||
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
|
||||
config.pop("_diffusers_version")
|
||||
@@ -276,14 +273,11 @@ class VAELoader(ComponentLoader):
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
|
||||
vae = vae_cls(**config).to(inference_args.device)
|
||||
|
||||
|
||||
# Find all safetensors files
|
||||
safetensors_list = glob.glob(
|
||||
os.path.join(str(model_path), "*.safetensors"))
|
||||
safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors"))
|
||||
# TODO(PY)
|
||||
assert len(
|
||||
safetensors_list
|
||||
) == 1, f"Found {len(safetensors_list)} safetensors files in {model_path}"
|
||||
assert len(safetensors_list) == 1, f"Found {len(safetensors_list)} safetensors files in {d}"
|
||||
loaded = safetensors_load_file(safetensors_list[0])
|
||||
vae.load_state_dict(loaded)
|
||||
dtype = PRECISION_TO_TYPE[inference_args.vae_precision]
|
||||
@@ -296,124 +290,103 @@ class VAELoader(ComponentLoader):
|
||||
}
|
||||
|
||||
vae.kwargs = vae_kwargs
|
||||
|
||||
|
||||
return vae
|
||||
|
||||
|
||||
class TransformerLoader(ComponentLoader):
|
||||
"""Loader for transformer."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
|
||||
"""Load the transformer based on the model path, architecture, and inference args."""
|
||||
model_config = get_diffusers_config(model=model_path)
|
||||
cls_name = model_config.pop("_class_name")
|
||||
if cls_name is None:
|
||||
raise ValueError(
|
||||
"Model config does not contain a _class_name attribute. "
|
||||
"Only diffusers format is supported.")
|
||||
raise ValueError(f"Model config does not contain a _class_name attribute. "
|
||||
"Only diffusers format is supported.")
|
||||
model_config.pop("_diffusers_version")
|
||||
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
||||
|
||||
# Find all safetensors files
|
||||
safetensors_list = glob.glob(
|
||||
os.path.join(str(model_path), "*.safetensors"))
|
||||
safetensors_list = glob.glob(os.path.join(str(model_path), "*.safetensors"))
|
||||
if not safetensors_list:
|
||||
raise ValueError(f"No safetensors files found in {model_path}")
|
||||
|
||||
logger.info("Loading model from %s safetensors files in %s",
|
||||
len(safetensors_list), model_path)
|
||||
|
||||
logger.info(f"Loading model from {len(safetensors_list)} safetensors files in {model_path}")
|
||||
|
||||
# initialize_sequence_parallel_group(inference_args.sp_size)
|
||||
|
||||
|
||||
# Load the model using FSDP loader
|
||||
logger.info("Loading model from %s", cls_name)
|
||||
model = load_fsdp_model(model_cls=model_cls,
|
||||
init_params=model_config,
|
||||
weight_dir_list=safetensors_list,
|
||||
device=inference_args.device,
|
||||
cpu_offload=inference_args.use_cpu_offload)
|
||||
|
||||
logger.info(f"Loading model from {cls_name}")
|
||||
model = load_fsdp_model(
|
||||
model_cls=model_cls,
|
||||
init_params=model_config,
|
||||
weight_dir_list=safetensors_list,
|
||||
device=inference_args.device,
|
||||
cpu_offload=inference_args.use_cpu_offload
|
||||
)
|
||||
|
||||
total_params = sum(p.numel() for p in model.parameters())
|
||||
logger.info("Loaded model with %.2fB parameters", total_params / 1e9)
|
||||
|
||||
logger.info(f"Loaded model with {total_params / 1e9:.2f}B parameters")
|
||||
|
||||
model.eval()
|
||||
return model
|
||||
|
||||
|
||||
class SchedulerLoader(ComponentLoader):
|
||||
"""Loader for scheduler."""
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
|
||||
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
|
||||
"""Load the scheduler based on the model path, architecture, and inference args."""
|
||||
if hasattr(inference_args,
|
||||
'denoise_type') and inference_args.denoise_type == "flow":
|
||||
# TODO(will): add schedulers to register or create a new scheduler registry
|
||||
# TODO(will): default to config file but allow override through
|
||||
# inference args. Currently only uses inference args.
|
||||
from fastvideo.v1.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchDiscreteScheduler)
|
||||
scheduler = FlowMatchDiscreteScheduler(
|
||||
shift=inference_args.flow_shift,
|
||||
solver=inference_args.flow_solver,
|
||||
)
|
||||
logger.info("Scheduler loaded: %s", scheduler)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid denoise type: {inference_args.denoise_type}")
|
||||
|
||||
|
||||
scheduler = get_scheduler(
|
||||
module_path=model_path,
|
||||
architecture=architecture,
|
||||
inference_args=inference_args,
|
||||
)
|
||||
logger.info(f"Scheduler loaded: {scheduler}")
|
||||
return scheduler
|
||||
|
||||
|
||||
class GenericComponentLoader(ComponentLoader):
|
||||
"""Generic loader for components that don't have a specific loader."""
|
||||
|
||||
def __init__(self, library="transformers") -> None:
|
||||
|
||||
def __init__(self, library="transformers"):
|
||||
super().__init__()
|
||||
self.library = library
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
|
||||
def load(self, model_path: str, architecture: str, inference_args: InferenceArgs):
|
||||
"""Load a generic component based on the model path, architecture, and inference args."""
|
||||
logger.warning("Using generic loader for %s with library %s",
|
||||
model_path, self.library)
|
||||
|
||||
logger.warning(f"Using generic loader for {model_path} with library {self.library}")
|
||||
|
||||
if self.library == "transformers":
|
||||
from transformers import AutoModel
|
||||
|
||||
|
||||
model = AutoModel.from_pretrained(
|
||||
model_path,
|
||||
trust_remote_code=inference_args.trust_remote_code,
|
||||
revision=inference_args.revision,
|
||||
)
|
||||
logger.info("Loaded generic transformers model: %s",
|
||||
model.__class__.__name__)
|
||||
logger.info(f"Loaded generic transformers model: {model.__class__.__name__}")
|
||||
return model
|
||||
elif self.library == "diffusers":
|
||||
logger.warning(
|
||||
"Generic loading for diffusers components is not fully implemented"
|
||||
)
|
||||
|
||||
logger.warning(f"Generic loading for diffusers components is not fully implemented")
|
||||
from fastvideo.v1.models.hf_transformer_utils import get_diffusers_config
|
||||
|
||||
model_config = get_diffusers_config(model=model_path)
|
||||
logger.info("Diffusers Model config: %s", model_config)
|
||||
logger.info(f"Diffusers Model config: {model_config}")
|
||||
# This is a placeholder - in a real implementation, you'd need to handle this properly
|
||||
return None
|
||||
else:
|
||||
raise ValueError(f"Unsupported library: {self.library}")
|
||||
|
||||
|
||||
class PipelineComponentLoader:
|
||||
"""
|
||||
Utility class for loading pipeline components.
|
||||
This replaces the chain of if-else statements in load_pipeline_module.
|
||||
"""
|
||||
|
||||
|
||||
@staticmethod
|
||||
def load_module(module_name: str, component_model_path: str,
|
||||
transformers_or_diffusers: str, architecture: str,
|
||||
inference_args: InferenceArgs):
|
||||
def load_module(module_name: str, component_model_path: str, transformers_or_diffusers: str,
|
||||
architecture: str, inference_args: InferenceArgs):
|
||||
"""
|
||||
Load a pipeline module.
|
||||
|
||||
@@ -427,16 +400,10 @@ class PipelineComponentLoader:
|
||||
Returns:
|
||||
The loaded module
|
||||
"""
|
||||
logger.info(
|
||||
"Loading %s using %s from %s",
|
||||
module_name,
|
||||
transformers_or_diffusers,
|
||||
component_model_path,
|
||||
)
|
||||
|
||||
logger.info(f"Loading {module_name} using {transformers_or_diffusers} from {component_model_path}")
|
||||
|
||||
# Get the appropriate loader for this module type
|
||||
loader = ComponentLoader.for_module_type(module_name,
|
||||
transformers_or_diffusers)
|
||||
|
||||
loader = ComponentLoader.for_module_type(module_name, transformers_or_diffusers)
|
||||
|
||||
# Load the module
|
||||
return loader.load(component_model_path, architecture, inference_args)
|
||||
|
||||
@@ -1,27 +1,21 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Adapted from torchtune
|
||||
# Copyright 2024 The TorchTune Authors.
|
||||
# Copyright 2025 The FastVideo Authors.
|
||||
|
||||
import contextlib
|
||||
import re
|
||||
from collections import defaultdict
|
||||
from typing import Any, Callable, Dict, Generator, List, Optional, Tuple, Type
|
||||
from itertools import chain
|
||||
from typing import (Any, Callable, DefaultDict, Dict, Generator, Hashable, List,
|
||||
Optional, Tuple, Type)
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from torch.distributed import DeviceMesh, init_device_mesh
|
||||
from fastvideo.v1.distributed.parallel_state import get_sequence_model_parallel_world_size
|
||||
from torch.distributed._composable.fsdp import CPUOffloadPolicy, fully_shard
|
||||
from torch.distributed._tensor import distribute_tensor
|
||||
from torch.nn.modules.module import _IncompatibleKeys
|
||||
from vllm.model_executor.model_loader.weight_utils import safetensors_weights_iterator
|
||||
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
get_sequence_model_parallel_world_size)
|
||||
from fastvideo.v1.models.loader.weight_utils import safetensors_weights_iterator
|
||||
|
||||
import contextlib
|
||||
import re
|
||||
|
||||
# TODO(PY): move this to utils elsewhere
|
||||
@contextlib.contextmanager
|
||||
@@ -51,8 +45,7 @@ def set_default_dtype(dtype: torch.dtype) -> Generator[None, None, None]:
|
||||
torch.set_default_dtype(old_dtype)
|
||||
|
||||
|
||||
def get_param_names_mapping(
|
||||
mapping_dict: Dict[str, str]) -> Callable[[str], tuple[str, Any, Any]]:
|
||||
def get_param_names_mapping(mapping_dict: Dict[str, str]) -> Callable[[str], str]:
|
||||
"""
|
||||
Creates a mapping function that transforms parameter names using regex patterns.
|
||||
|
||||
@@ -63,9 +56,8 @@ def get_param_names_mapping(
|
||||
Returns:
|
||||
Callable[[str], str]: A function that maps parameter names from source to target format
|
||||
"""
|
||||
|
||||
def mapping_fn(name: str) -> tuple[str, Any, Any]:
|
||||
|
||||
def mapping_fn(name: str) -> str:
|
||||
|
||||
# Try to match and transform the name using the regex patterns in mapping_dict
|
||||
for pattern, replacement in mapping_dict.items():
|
||||
match = re.match(pattern, name)
|
||||
@@ -75,13 +67,13 @@ def get_param_names_mapping(
|
||||
if isinstance(replacement, tuple):
|
||||
merge_index = replacement[1]
|
||||
total_splitted_params = replacement[2]
|
||||
replacement = replacement[0]
|
||||
replacement= replacement[0]
|
||||
name = re.sub(pattern, replacement, name)
|
||||
return name, merge_index, total_splitted_params
|
||||
|
||||
|
||||
# If no pattern matches, return the original name
|
||||
return name, None, None
|
||||
|
||||
|
||||
return mapping_fn
|
||||
|
||||
|
||||
@@ -98,15 +90,12 @@ def load_fsdp_model(
|
||||
model = model_cls(**init_params)
|
||||
device_mesh = init_device_mesh(
|
||||
"cuda",
|
||||
mesh_shape=(get_sequence_model_parallel_world_size(), ),
|
||||
mesh_shape=(get_sequence_model_parallel_world_size(),),
|
||||
mesh_dim_names=("dp", ),
|
||||
)
|
||||
shard_model(model,
|
||||
cpu_offload=cpu_offload,
|
||||
reshard_after_forward=True,
|
||||
dp_mesh=device_mesh["dp"])
|
||||
shard_model(model, cpu_offload=cpu_offload, reshard_after_forward=True, dp_mesh=device_mesh["dp"])
|
||||
weight_iterator = safetensors_weights_iterator(weight_dir_list)
|
||||
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
|
||||
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
|
||||
load_fsdp_model_from_full_model_state_dict(
|
||||
model,
|
||||
weight_iterator,
|
||||
@@ -117,13 +106,11 @@ def load_fsdp_model(
|
||||
)
|
||||
for n, p in chain(model.named_parameters(), model.named_buffers()):
|
||||
if p.is_meta:
|
||||
raise RuntimeError(
|
||||
f"Unexpected param or buffer {n} on meta device.")
|
||||
raise RuntimeError(f"Unexpected param or buffer {n} on meta device.")
|
||||
for p in model.parameters():
|
||||
p.requires_grad = False
|
||||
p.requires_grad = False
|
||||
return model
|
||||
|
||||
|
||||
def shard_model(
|
||||
model,
|
||||
*,
|
||||
@@ -148,16 +135,13 @@ def shard_model(
|
||||
reshard_after_forward (bool): Whether to reshard parameters and buffers after
|
||||
the forward pass. Setting this to True corresponds to the FULL_SHARD sharding strategy
|
||||
from FSDP1, while setting it to False corresponds to the SHARD_GRAD_OP sharding strategy.
|
||||
dp_mesh (Optional[DeviceMesh]): Device mesh to use for FSDP sharding under multiple parallelism.
|
||||
dp_mesh (Optional[DeviceMesh]): Device mesh to use for FSDP sharding under mutliple parallelism.
|
||||
Default to None.
|
||||
|
||||
Raises:
|
||||
ValueError: If no layer modules were sharded, indicating that no shard_condition was triggered.
|
||||
"""
|
||||
fsdp_kwargs = {
|
||||
"reshard_after_forward": reshard_after_forward,
|
||||
"mesh": dp_mesh
|
||||
}
|
||||
fsdp_kwargs = {"reshard_after_forward": reshard_after_forward, "mesh": dp_mesh}
|
||||
if cpu_offload:
|
||||
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy()
|
||||
|
||||
@@ -165,10 +149,7 @@ def shard_model(
|
||||
# lowest-level modules first
|
||||
num_layers_sharded = 0
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
if any([
|
||||
shard_condition(n, m)
|
||||
for shard_condition in model._fsdp_shard_conditions
|
||||
]):
|
||||
if any([shard_condition(n, m) for shard_condition in model._fsdp_shard_conditions]):
|
||||
fully_shard(m, **fsdp_kwargs)
|
||||
num_layers_sharded += 1
|
||||
|
||||
@@ -179,8 +160,7 @@ def shard_model(
|
||||
|
||||
# Finally shard the entire model to account for any stragglers
|
||||
fully_shard(model, **fsdp_kwargs)
|
||||
|
||||
|
||||
|
||||
# TODO(PY): device mesh for cfg parallel
|
||||
def load_fsdp_model_from_full_model_state_dict(
|
||||
model: torch.nn.Module,
|
||||
@@ -188,7 +168,7 @@ def load_fsdp_model_from_full_model_state_dict(
|
||||
device: torch.device,
|
||||
strict: bool = False,
|
||||
cpu_offload: bool = False,
|
||||
param_names_mapping: Optional[Callable[[str], tuple[str, Any, Any]]] = None,
|
||||
param_names_mapping: Optional[Callable[[str], str]] = None,
|
||||
) -> _IncompatibleKeys:
|
||||
"""
|
||||
Converting full state dict into a sharded state dict
|
||||
@@ -210,34 +190,27 @@ def load_fsdp_model_from_full_model_state_dict(
|
||||
NotImplementedError: If got FSDP with more than 1D.
|
||||
"""
|
||||
meta_sharded_sd = model.state_dict()
|
||||
|
||||
|
||||
sharded_sd = {}
|
||||
to_merge_params: DefaultDict[Hashable, Dict[Any, Any]] = defaultdict(dict)
|
||||
to_merge_params = defaultdict(dict)
|
||||
for source_param_name, full_tensor in full_sd_iterator:
|
||||
assert param_names_mapping is not None
|
||||
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
|
||||
source_param_name)
|
||||
|
||||
target_param_name, merge_index, num_params_to_merge = param_names_mapping(source_param_name)
|
||||
|
||||
if merge_index is not None:
|
||||
to_merge_params[target_param_name][merge_index] = full_tensor
|
||||
if len(to_merge_params[target_param_name]) == num_params_to_merge:
|
||||
# cat at dim=1 according to the merge_index order
|
||||
sorted_tensors = [
|
||||
to_merge_params[target_param_name][i]
|
||||
for i in range(num_params_to_merge)
|
||||
]
|
||||
sorted_tensors = [to_merge_params[target_param_name][i] for i in range(num_params_to_merge)]
|
||||
full_tensor = torch.cat(sorted_tensors, dim=0)
|
||||
del to_merge_params[target_param_name]
|
||||
else:
|
||||
continue
|
||||
|
||||
|
||||
sharded_meta_param = meta_sharded_sd.get(target_param_name)
|
||||
if sharded_meta_param is None:
|
||||
raise ValueError(
|
||||
f"Parameter {source_param_name}-->{target_param_name} not found in meta sharded state dict"
|
||||
)
|
||||
raise ValueError(f"Parameter {source_param_name}-->{target_param_name} not found in meta sharded state dict")
|
||||
full_tensor = full_tensor.to(sharded_meta_param.dtype).to(device)
|
||||
|
||||
|
||||
if not hasattr(sharded_meta_param, "device_mesh"):
|
||||
# In cases where parts of the model aren't sharded, some parameters will be plain tensors
|
||||
sharded_tensor = full_tensor
|
||||
@@ -251,4 +224,4 @@ def load_fsdp_model_from_full_model_state_dict(
|
||||
sharded_tensor = sharded_tensor.cpu()
|
||||
sharded_sd[target_param_name] = nn.Parameter(sharded_tensor)
|
||||
# choose `assign=True` since we cannot call `copy_` on meta tensor
|
||||
return model.load_state_dict(sharded_sd, strict=strict, assign=True)
|
||||
return model.load_state_dict(sharded_sd, strict=strict, assign=True)
|
||||
@@ -1,7 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Utilities for selecting and loading models."""
|
||||
import contextlib
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
@@ -1,5 +1,9 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/model_loader/weight_utils.py
|
||||
|
||||
# Adapted from vllm
|
||||
# Copyright 2023 The vLLM Authors.
|
||||
# Copyright 2025 The FastVideo Authors.
|
||||
|
||||
"""Utilities for downloading and initializing model weights."""
|
||||
import fnmatch
|
||||
import hashlib
|
||||
@@ -29,7 +33,7 @@ logger = init_logger(__name__)
|
||||
temp_dir = tempfile.gettempdir()
|
||||
|
||||
|
||||
def enable_hf_transfer() -> None:
|
||||
def enable_hf_transfer():
|
||||
"""automatically activates hf_transfer
|
||||
"""
|
||||
if "HF_HUB_ENABLE_HF_TRANSFER" not in os.environ:
|
||||
@@ -60,7 +64,8 @@ def get_lock(model_name_or_path: Union[str, Path],
|
||||
# add hash to avoid conflict with old users' lock files
|
||||
lock_file_name = hash_name + model_name + ".lock"
|
||||
# mode 0o666 is required for the filelock to be shared across users
|
||||
lock = filelock.FileLock(os.path.join(lock_dir, lock_file_name), mode=0o666)
|
||||
lock = filelock.FileLock(os.path.join(lock_dir, lock_file_name),
|
||||
mode=0o666)
|
||||
return lock
|
||||
|
||||
|
||||
@@ -117,7 +122,7 @@ def download_weights_from_hf(
|
||||
# downloading the same model weights at the same time.
|
||||
with get_lock(model_name_or_path, cache_dir):
|
||||
start_time = time.perf_counter()
|
||||
hf_folder: str = snapshot_download(
|
||||
hf_folder = snapshot_download(
|
||||
model_name_or_path,
|
||||
allow_patterns=allow_patterns,
|
||||
ignore_patterns=ignore_patterns,
|
||||
@@ -338,4 +343,4 @@ def maybe_remap_kv_scale_name(name: str, params_dict: dict) -> Optional[str]:
|
||||
return remapped_name
|
||||
|
||||
# If there were no matches, return the untouched param name
|
||||
return name
|
||||
return name
|
||||
@@ -2,7 +2,7 @@
|
||||
# Adapted from: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/parameter.py
|
||||
|
||||
from fractions import Fraction
|
||||
from typing import Any, Callable, Optional, Tuple, Union
|
||||
from typing import Callable, Optional, Union
|
||||
|
||||
import torch
|
||||
from torch.nn import Parameter
|
||||
@@ -11,6 +11,12 @@ from fastvideo.v1.distributed import get_tensor_model_parallel_rank
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.utils import _make_synced_weight_loader
|
||||
|
||||
__all__ = [
|
||||
"BasevLLMParameter", "PackedvLLMParameter", "PerTensorScaleParameter",
|
||||
"ModelWeightParameter", "ChannelQuantScaleParameter",
|
||||
"GroupQuantScaleParameter", "PackedColumnParameter", "RowvLLMParameter"
|
||||
]
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -43,7 +49,7 @@ class BasevLLMParameter(Parameter):
|
||||
# tensor, which is param.data, leading to the redundant memory usage.
|
||||
# This sometimes causes OOM errors during model loading. To avoid this,
|
||||
# we sync the param tensor after its weight loader is called.
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
from vllm.platforms import current_platform
|
||||
if current_platform.is_tpu():
|
||||
weight_loader = _make_synced_weight_loader(weight_loader)
|
||||
|
||||
@@ -295,8 +301,7 @@ class PackedColumnParameter(_ColumnvLLMParameter):
|
||||
def marlin_tile_size(self):
|
||||
return self._marlin_tile_size
|
||||
|
||||
def adjust_shard_indexes_for_packing(self, shard_size,
|
||||
shard_offset) -> Tuple[Any, Any]:
|
||||
def adjust_shard_indexes_for_packing(self, shard_size, shard_offset):
|
||||
return _adjust_shard_indexes_for_packing(
|
||||
shard_size=shard_size,
|
||||
shard_offset=shard_offset,
|
||||
@@ -413,12 +418,12 @@ def permute_param_layout_(param: BasevLLMParameter, input_dim: int,
|
||||
|
||||
|
||||
def _adjust_shard_indexes_for_marlin(shard_size, shard_offset,
|
||||
marlin_tile_size) -> Tuple[Any, Any]:
|
||||
marlin_tile_size):
|
||||
return shard_size * marlin_tile_size, shard_offset * marlin_tile_size
|
||||
|
||||
|
||||
def _adjust_shard_indexes_for_packing(shard_size, shard_offset, packed_factor,
|
||||
marlin_tile_size) -> Tuple[Any, Any]:
|
||||
marlin_tile_size):
|
||||
shard_size = shard_size // packed_factor
|
||||
shard_offset = shard_offset // packed_factor
|
||||
if marlin_tile_size is not None:
|
||||
@@ -426,4 +431,4 @@ def _adjust_shard_indexes_for_packing(shard_size, shard_offset, packed_factor,
|
||||
shard_size=shard_size,
|
||||
shard_offset=shard_offset,
|
||||
marlin_tile_size=marlin_tile_size)
|
||||
return shard_size, shard_offset
|
||||
return shard_size, shard_offset
|
||||
@@ -1,27 +1,21 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/models/registry.py
|
||||
|
||||
import importlib
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
import os
|
||||
import pickle
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import AbstractSet, Callable, Dict, List, Optional, Tuple, Type, Union, TypeVar
|
||||
import importlib
|
||||
from functools import lru_cache
|
||||
from typing import (AbstractSet, Callable, Dict, List, NoReturn, Optional,
|
||||
Tuple, Type, TypeVar, Union, cast)
|
||||
|
||||
import cloudpickle
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.v1.logger import logger
|
||||
|
||||
# huggingface class name: (component_name, fastvideo module name, fastvideo class name)
|
||||
_TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"HunyuanVideoTransformer3DModel":
|
||||
("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
|
||||
"HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
|
||||
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
|
||||
}
|
||||
|
||||
@@ -32,18 +26,15 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
|
||||
_TEXT_ENCODER_MODELS = {
|
||||
"CLIPTextModel": ("encoders", "clip", "CLIPTextModel"),
|
||||
"LlamaModel": ("encoders", "llama", "LlamaModel"),
|
||||
"UMT5EncoderModel": ("encoders", "t5", "UMT5EncoderModel"),
|
||||
"LlamaModel": ("encoders", "llama", "LlamaModel"),
|
||||
}
|
||||
|
||||
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
|
||||
_IMAGE_ENCODER_MODELS = {
|
||||
# "HunyuanVideoTransformer3DModel": ("image_encoder", "hunyuanvideo", "HunyuanVideoImageEncoder"),
|
||||
}
|
||||
|
||||
_VAE_MODELS = {
|
||||
"AutoencoderKLHunyuanVideo":
|
||||
("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
|
||||
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
|
||||
"AutoencoderKLHunyuanVideo": ("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
|
||||
}
|
||||
|
||||
_FAST_VIDEO_MODELS = {
|
||||
@@ -58,16 +49,18 @@ _SUBPROCESS_COMMAND = [
|
||||
sys.executable, "-m", "fastvideo.v1.models.dits.registry"
|
||||
]
|
||||
|
||||
|
||||
_T = TypeVar("_T")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _ModelInfo:
|
||||
architecture: str
|
||||
|
||||
|
||||
@staticmethod
|
||||
def from_model_cls(model: Type[nn.Module]) -> "_ModelInfo":
|
||||
return _ModelInfo(architecture=model.__name__, )
|
||||
return _ModelInfo(
|
||||
architecture=model.__name__,)
|
||||
|
||||
|
||||
class _BaseRegisteredModel(ABC):
|
||||
@@ -103,7 +96,6 @@ class _RegisteredModel(_BaseRegisteredModel):
|
||||
def load_model_cls(self) -> Type[nn.Module]:
|
||||
return self.model_cls
|
||||
|
||||
|
||||
def _run_in_subprocess(fn: Callable[[], _T]) -> _T:
|
||||
# NOTE: We use a temporary directory instead of a temporary file to avoid
|
||||
# issues like https://stackoverflow.com/questions/23212435/permission-denied-to-write-to-my-temporary-file
|
||||
@@ -128,9 +120,9 @@ def _run_in_subprocess(fn: Callable[[], _T]) -> _T:
|
||||
f"{returned.stderr.decode()}") from e
|
||||
|
||||
with open(output_filepath, "rb") as f:
|
||||
return cast(_T, pickle.load(f))
|
||||
|
||||
|
||||
return pickle.load(f)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _LazyRegisteredModel(_BaseRegisteredModel):
|
||||
"""
|
||||
@@ -147,7 +139,7 @@ class _LazyRegisteredModel(_BaseRegisteredModel):
|
||||
|
||||
def load_model_cls(self) -> Type[nn.Module]:
|
||||
mod = importlib.import_module(self.module_name)
|
||||
return cast(Type[nn.Module], getattr(mod, self.class_name))
|
||||
return getattr(mod, self.class_name)
|
||||
|
||||
|
||||
@lru_cache(maxsize=128)
|
||||
@@ -160,7 +152,8 @@ def _try_load_model_cls(
|
||||
try:
|
||||
return model.load_model_cls()
|
||||
except Exception:
|
||||
logger.exception("Error in loading model architecture '%s'", model_arch)
|
||||
logger.exception("Error in loading model architecture '%s'",
|
||||
model_arch)
|
||||
return None
|
||||
|
||||
|
||||
@@ -204,7 +197,8 @@ class _ModelRegistry:
|
||||
if model_arch in self.models:
|
||||
logger.warning(
|
||||
"Model architecture %s is already registered, and will be "
|
||||
"overwritten by the new model class %s.", model_arch, model_cls)
|
||||
"overwritten by the new model class %s.", model_arch,
|
||||
model_cls)
|
||||
|
||||
if isinstance(model_cls, str):
|
||||
split_str = model_cls.split(":")
|
||||
@@ -218,7 +212,7 @@ class _ModelRegistry:
|
||||
|
||||
self.models[model_arch] = model
|
||||
|
||||
def _raise_for_unsupported(self, architectures: List[str]) -> NoReturn:
|
||||
def _raise_for_unsupported(self, architectures: List[str]):
|
||||
all_supported_archs = self.get_supported_archs()
|
||||
|
||||
if any(arch in all_supported_archs for arch in architectures):
|
||||
@@ -230,7 +224,8 @@ class _ModelRegistry:
|
||||
f"Model architectures {architectures} are not supported for now. "
|
||||
f"Supported architectures: {all_supported_archs}")
|
||||
|
||||
def _try_load_model_cls(self, model_arch: str) -> Optional[Type[nn.Module]]:
|
||||
def _try_load_model_cls(self,
|
||||
model_arch: str) -> Optional[Type[nn.Module]]:
|
||||
if model_arch not in self.models:
|
||||
return None
|
||||
|
||||
@@ -284,6 +279,8 @@ class _ModelRegistry:
|
||||
|
||||
return self._raise_for_unsupported(architectures)
|
||||
|
||||
|
||||
|
||||
|
||||
ModelRegistry = _ModelRegistry({
|
||||
model_arch:
|
||||
@@ -292,6 +289,5 @@ ModelRegistry = _ModelRegistry({
|
||||
component_name=component_name,
|
||||
class_name=cls_name,
|
||||
)
|
||||
for model_arch, (component_name, mod_relname,
|
||||
cls_name) in _FAST_VIDEO_MODELS.items()
|
||||
})
|
||||
for model_arch, (component_name, mod_relname, cls_name) in _FAST_VIDEO_MODELS.items()
|
||||
})
|
||||
@@ -1,5 +1,3 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# Copyright 2024 Stability AI, Katherine Crowson and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
@@ -20,7 +18,7 @@
|
||||
# ==============================================================================
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Optional, Tuple, Union
|
||||
from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
@@ -63,7 +61,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
Whether to reverse the timestep schedule.
|
||||
"""
|
||||
|
||||
_compatibles: list[Any] = []
|
||||
_compatibles = []
|
||||
order = 1
|
||||
|
||||
@register_to_config
|
||||
@@ -82,17 +80,14 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
|
||||
self.sigmas = sigmas
|
||||
# the value fed to model
|
||||
self.timesteps = (sigmas[:-1] *
|
||||
num_train_timesteps).to(dtype=torch.float32)
|
||||
self.timesteps = (sigmas[:-1] * num_train_timesteps).to(dtype=torch.float32)
|
||||
|
||||
self._step_index: int | None = None
|
||||
self._begin_index = 0
|
||||
self._step_index = None
|
||||
self._begin_index = None
|
||||
|
||||
self.supported_solver = ["euler"]
|
||||
if solver not in self.supported_solver:
|
||||
raise ValueError(
|
||||
f"Solver {solver} not supported. Supported solvers: {self.supported_solver}"
|
||||
)
|
||||
raise ValueError(f"Solver {solver} not supported. Supported solvers: {self.supported_solver}")
|
||||
|
||||
@property
|
||||
def step_index(self):
|
||||
@@ -126,7 +121,7 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
self,
|
||||
num_inference_steps: int,
|
||||
device: Union[str, torch.device] = None,
|
||||
n_tokens: int = 0,
|
||||
n_tokens: int = None,
|
||||
):
|
||||
"""
|
||||
Sets the discrete timesteps used for the diffusion chain (to be run before inference).
|
||||
@@ -148,13 +143,12 @@ 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
|
||||
|
||||
def index_for_timestep(self, timestep, schedule_timesteps=None) -> int:
|
||||
def index_for_timestep(self, timestep, schedule_timesteps=None):
|
||||
if schedule_timesteps is None:
|
||||
schedule_timesteps = self.timesteps
|
||||
|
||||
@@ -166,11 +160,9 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
# case we start in the middle of the denoising schedule (e.g. for image-to-image)
|
||||
pos = 1 if len(indices) > 1 else 0
|
||||
|
||||
idx: int = indices[pos].item()
|
||||
return indices[pos].item()
|
||||
|
||||
return idx
|
||||
|
||||
def _init_step_index(self, timestep) -> None:
|
||||
def _init_step_index(self, timestep):
|
||||
if self.begin_index is None:
|
||||
if isinstance(timestep, torch.Tensor):
|
||||
timestep = timestep.to(self.timesteps.device)
|
||||
@@ -178,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):
|
||||
@@ -218,11 +208,11 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
returned, otherwise a tuple is returned where the first element is the sample tensor.
|
||||
"""
|
||||
|
||||
if isinstance(timestep, (int, torch.IntTensor, 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)
|
||||
@@ -230,18 +220,14 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
# Upcast to avoid precision issues when computing prev_sample
|
||||
sample = sample.to(torch.float32)
|
||||
|
||||
assert self.step_index is not None
|
||||
dt = self.sigmas[self.step_index + 1] - self.sigmas[self.step_index]
|
||||
|
||||
if self.config.solver == "euler":
|
||||
prev_sample = sample + model_output.to(torch.float32) * dt
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}"
|
||||
)
|
||||
raise ValueError(f"Solver {self.config.solver} not supported. Supported solvers: {self.supported_solver}")
|
||||
|
||||
# upon completion increase step index by one
|
||||
assert self._step_index is not None
|
||||
self._step_index += 1
|
||||
|
||||
if not return_dict:
|
||||
@@ -250,4 +236,4 @@ class FlowMatchDiscreteScheduler(SchedulerMixin, ConfigMixin):
|
||||
return FlowMatchDiscreteSchedulerOutput(prev_sample=prev_sample)
|
||||
|
||||
def __len__(self):
|
||||
return self.config.num_train_timesteps
|
||||
return self.config.num_train_timesteps
|
||||
@@ -1,20 +1,17 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Optional
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from transformers.utils import ModelOutput
|
||||
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from transformers.utils import ModelOutput
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def use_default(value, default) -> Any:
|
||||
def use_default(value, default):
|
||||
return value if value is not None else default
|
||||
|
||||
|
||||
@dataclass
|
||||
class TextEncoderModelOutput(ModelOutput):
|
||||
"""
|
||||
@@ -23,7 +20,6 @@ class TextEncoderModelOutput(ModelOutput):
|
||||
Args:
|
||||
hidden_state (`torch.FloatTensor` of shape `(batch_size, sequence_length, hidden_size)`):
|
||||
Sequence of hidden-states at the output of the last layer of the model.
|
||||
|
||||
attention_mask (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*):
|
||||
Mask to avoid performing attention on padding token indices. Mask values selected in ``[0, 1]``:
|
||||
hidden_states_list (`tuple(torch.FloatTensor)`, *optional*, returned when `output_hidden_states=True` is passed):
|
||||
@@ -64,13 +60,13 @@ class TextEncoder(nn.Module):
|
||||
self.model_path = text_encoder_path
|
||||
self.use_attention_mask = use_attention_mask
|
||||
if prompt_template_video is not None:
|
||||
assert (use_attention_mask is True
|
||||
), "Attention mask is True required when training videos."
|
||||
assert (use_attention_mask is True), "Attention mask is True required when training videos."
|
||||
self.prompt_template = prompt_template
|
||||
self.prompt_template_video = prompt_template_video
|
||||
self.hidden_state_skip_layer = hidden_state_skip_layer
|
||||
self.apply_final_norm = apply_final_norm
|
||||
|
||||
|
||||
if "T5" in self.text_encoder_type:
|
||||
self.output_key = output_key or "last_hidden_state"
|
||||
elif "CLIPTextModel" in self.text_encoder_type:
|
||||
@@ -78,9 +74,8 @@ class TextEncoder(nn.Module):
|
||||
elif "LlamaModel" in self.text_encoder_type or "glm" in self.text_encoder_type:
|
||||
self.output_key = output_key or "last_hidden_state"
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported text encoder type: {self.text_encoder_type}")
|
||||
|
||||
raise ValueError(f"Unsupported text encoder type: {self.text_encoder_type}")
|
||||
|
||||
self.model = text_encoder
|
||||
# self.dtype = self.model.dtype
|
||||
self.device = device
|
||||
@@ -91,7 +86,7 @@ class TextEncoder(nn.Module):
|
||||
return f"{self.text_encoder_type} ({self.precision} - {self.model_path})"
|
||||
|
||||
@staticmethod
|
||||
def apply_text_to_template(text, template, prevent_empty_text=True) -> str:
|
||||
def apply_text_to_template(text, template, prevent_empty_text=True):
|
||||
"""
|
||||
Apply text to template.
|
||||
|
||||
@@ -107,7 +102,7 @@ class TextEncoder(nn.Module):
|
||||
else:
|
||||
raise TypeError(f"Unsupported template type: {type(template)}")
|
||||
|
||||
def text2tokens(self, text) -> dict:
|
||||
def text2tokens(self, text):
|
||||
"""
|
||||
Tokenize the input text.
|
||||
|
||||
@@ -116,22 +111,23 @@ class TextEncoder(nn.Module):
|
||||
"""
|
||||
if self.prompt_template_video is not None:
|
||||
prompt_template = self.prompt_template_video["template"]
|
||||
|
||||
|
||||
text = self.apply_text_to_template(text, prompt_template)
|
||||
|
||||
|
||||
kwargs = dict(
|
||||
truncation=True,
|
||||
max_length=self.max_length,
|
||||
padding="max_length",
|
||||
return_tensors="pt",
|
||||
)
|
||||
batch_encoding: dict = self.tokenizer(
|
||||
return self.tokenizer(
|
||||
text,
|
||||
return_length=False,
|
||||
return_overflowing_tokens=False,
|
||||
return_attention_mask=True,
|
||||
**kwargs,
|
||||
)
|
||||
return batch_encoding
|
||||
|
||||
def encode(
|
||||
self,
|
||||
@@ -139,7 +135,7 @@ class TextEncoder(nn.Module):
|
||||
use_attention_mask=None,
|
||||
hidden_state_skip_layer=None,
|
||||
device=None,
|
||||
) -> TextEncoderModelOutput:
|
||||
):
|
||||
"""
|
||||
Args:
|
||||
batch_encoding (dict): Batch encoding from tokenizer.
|
||||
@@ -153,37 +149,34 @@ class TextEncoder(nn.Module):
|
||||
return_texts (bool): Whether to return the decoded texts. Defaults to False.
|
||||
"""
|
||||
device = self.model.device if device is None else device
|
||||
use_attention_mask = use_default(use_attention_mask,
|
||||
self.use_attention_mask)
|
||||
hidden_state_skip_layer = use_default(hidden_state_skip_layer,
|
||||
self.hidden_state_skip_layer)
|
||||
use_attention_mask = use_default(use_attention_mask, self.use_attention_mask)
|
||||
hidden_state_skip_layer = use_default(hidden_state_skip_layer, self.hidden_state_skip_layer)
|
||||
attention_mask = (batch_encoding["attention_mask"].to(device) if use_attention_mask else None)
|
||||
|
||||
# note: clip will need attention mask
|
||||
# TODO(will): unify interface with dit
|
||||
# TODO (peiyuan): why clip need attention mask?
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs = self.model(
|
||||
input_ids=batch_encoding["input_ids"].to(device),
|
||||
output_hidden_states=hidden_state_skip_layer is not None,
|
||||
)
|
||||
outputs = self.model(
|
||||
input_ids=batch_encoding["input_ids"].to(device),
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=hidden_state_skip_layer is not None,
|
||||
)
|
||||
if hidden_state_skip_layer is not None:
|
||||
last_hidden_state = outputs.hidden_states[-(
|
||||
hidden_state_skip_layer + 1)]
|
||||
last_hidden_state = outputs.hidden_states[-(hidden_state_skip_layer + 1)]
|
||||
# Real last hidden state already has layer norm applied. So here we only apply it
|
||||
# for intermediate layers.
|
||||
if hidden_state_skip_layer > 0 and self.apply_final_norm:
|
||||
last_hidden_state = self.model.final_layer_norm(
|
||||
last_hidden_state)
|
||||
last_hidden_state = self.model.final_layer_norm(last_hidden_state)
|
||||
else:
|
||||
last_hidden_state = outputs[self.output_key]
|
||||
|
||||
# Remove hidden states of instruction tokens, only keep prompt tokens.
|
||||
if self.prompt_template_video is not None:
|
||||
|
||||
|
||||
crop_start = self.prompt_template_video.get("crop_start", -1)
|
||||
|
||||
last_hidden_state = last_hidden_state[:, crop_start:]
|
||||
|
||||
attention_mask = (attention_mask[:, crop_start:] if use_attention_mask else None)
|
||||
total_length = attention_mask.sum()
|
||||
last_hidden_state = last_hidden_state[:, :total_length]
|
||||
return TextEncoderModelOutput(last_hidden_state)
|
||||
|
||||
def forward(
|
||||
@@ -198,5 +191,7 @@ class TextEncoder(nn.Module):
|
||||
return self.encode(
|
||||
batch_encoding,
|
||||
use_attention_mask=use_attention_mask,
|
||||
output_hidden_states=output_hidden_states,
|
||||
hidden_state_skip_layer=hidden_state_skip_layer,
|
||||
return_texts=return_texts,
|
||||
)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user