Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e9f0f714ed | ||
|
|
554a8f045e | ||
|
|
6d0eb1b956 | ||
|
|
20f5d1a779 | ||
|
|
5da1c9aadf |
@@ -1,201 +0,0 @@
|
||||
env:
|
||||
IMAGE_VERSION: "py3.12-latest"
|
||||
BUILDKITE_CLEAN_CHECKOUT: true
|
||||
|
||||
steps:
|
||||
- label: "pre-commit"
|
||||
command: ".buildkite/scripts/pre_commit.sh"
|
||||
agents:
|
||||
queue: "default"
|
||||
|
||||
- wait
|
||||
|
||||
- label: "Trigger Tests"
|
||||
plugins:
|
||||
- monorepo-diff#v1.4.0:
|
||||
diff: 'git fetch origin "$BUILDKITE_PULL_REQUEST_BASE_BRANCH" && git diff --name-only origin/"$BUILDKITE_PULL_REQUEST_BASE_BRANCH"...HEAD'
|
||||
watch:
|
||||
- path:
|
||||
- "fastvideo/models/encoders/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/encoders/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Encoder Tests"
|
||||
env:
|
||||
- TEST_TYPE=encoder
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/models/vaes/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/vaes/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "VAE Tests"
|
||||
env:
|
||||
- TEST_TYPE=vae
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/models/dits/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/transformers/**"
|
||||
- "fastvideo/layers/**"
|
||||
- "fastvideo/attention/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Transformer Tests"
|
||||
env:
|
||||
- TEST_TYPE=transformer
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**/*.py"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 45m .buildkite/scripts/pr_test.sh"
|
||||
label: "SSIM Tests"
|
||||
env:
|
||||
- TEST_TYPE=ssim
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/tests/lora/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/transformers/**"
|
||||
- "fastvideo/pipelines/**"
|
||||
- "fastvideo/layers/lora/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "LoRA Inference Tests"
|
||||
env:
|
||||
- TEST_TYPE=inference_lora
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/training/*distillation_pipeline.py"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Distillation DMDTests"
|
||||
env:
|
||||
- TEST_TYPE=distillation_dmd
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "LoRA Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training_lora
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "csrc/attn/video_sparse_attn/**"
|
||||
- "csrc/attn/video_sparse_attn/tk/**"
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
- "csrc/attn/video_sparse_attn/config_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/vsa.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Training Tests VSA"
|
||||
env:
|
||||
- TEST_TYPE=training_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "csrc/attn/sliding_tile_attn/**"
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
- "csrc/attn/sliding_tile_attn/config_sta.py"
|
||||
- "csrc/attn/sliding_tile_attn/st_attn.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Inference Tests STA"
|
||||
env:
|
||||
- TEST_TYPE=inference_sta
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/sliding_tile_attn/**"
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
- "csrc/attn/sliding_tile_attn/config_sta.py"
|
||||
- "csrc/attn/sliding_tile_attn/st_attn.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests STA"
|
||||
env:
|
||||
- TEST_TYPE=precision_sta
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/video_sparse_attn/**"
|
||||
- "csrc/attn/video_sparse_attn/tk/**"
|
||||
- "csrc/attn/tests/test_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
- "csrc/attn/video_sparse_attn/config_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/vsa.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests VSA"
|
||||
env:
|
||||
- TEST_TYPE=precision_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/vmoba_attn/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests VMoBA"
|
||||
env:
|
||||
- TEST_TYPE=precision_vmoba
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/vmoba_attn/vmoba/**"
|
||||
- "fastvideo/attention/backends/vmoba.py"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Inference Tests VMoBA"
|
||||
env:
|
||||
- TEST_TYPE=inference_vmoba
|
||||
agents:
|
||||
queue: "default"
|
||||
@@ -1,138 +0,0 @@
|
||||
#!/bin/bash
|
||||
set -uo pipefail
|
||||
|
||||
log() {
|
||||
echo "[$(date '+%Y-%m-%d %H:%M:%S')] $1"
|
||||
}
|
||||
|
||||
log "=== Starting Modal test execution ==="
|
||||
|
||||
# Change to the project directory
|
||||
cd "$(dirname "$0")/../.."
|
||||
PROJECT_ROOT=$(pwd)
|
||||
log "Project root: $PROJECT_ROOT"
|
||||
|
||||
# Install Modal if not available
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Modal not found, installing..."
|
||||
python3 -m pip install modal
|
||||
|
||||
# Verify installation
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Error: Failed to install modal. Please install it manually."
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
|
||||
log "modal version: $(python3 -m modal --version)"
|
||||
|
||||
# Set up Modal authentication using Buildkite secrets
|
||||
log "Setting up Modal authentication from Buildkite secrets..."
|
||||
MODAL_TOKEN_ID=$(buildkite-agent secret get modal_token_id)
|
||||
MODAL_TOKEN_SECRET=$(buildkite-agent secret get modal_token_secret)
|
||||
|
||||
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
|
||||
|
||||
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
|
||||
|
||||
if [ -n "$MODAL_TOKEN_ID" ] && [ -n "$MODAL_TOKEN_SECRET" ]; then
|
||||
log "Retrieved Modal credentials from Buildkite secrets"
|
||||
python3 -m modal token set --token-id "$MODAL_TOKEN_ID" --token-secret "$MODAL_TOKEN_SECRET" --profile buildkite-ci --activate --verify
|
||||
if [ $? -eq 0 ]; then
|
||||
log "Modal authentication successful"
|
||||
else
|
||||
log "Error: Failed to set Modal credentials"
|
||||
exit 1
|
||||
fi
|
||||
else
|
||||
log "Error: Could not retrieve Modal credentials from Buildkite secrets."
|
||||
log "Please ensure 'modal_token_id' and 'modal_token_secret' secrets are set in Buildkite."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
MODAL_TEST_FILE="fastvideo/tests/modal/pr_test.py"
|
||||
|
||||
if [ -z "${TEST_TYPE:-}" ]; then
|
||||
log "Error: TEST_TYPE environment variable is not set"
|
||||
exit 1
|
||||
fi
|
||||
log "Test type: $TEST_TYPE"
|
||||
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$BUILDKITE_PULL_REQUEST IMAGE_VERSION=$IMAGE_VERSION"
|
||||
|
||||
case "$TEST_TYPE" in
|
||||
"encoder")
|
||||
log "Running encoder tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
|
||||
;;
|
||||
"vae")
|
||||
log "Running VAE tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
|
||||
;;
|
||||
"transformer")
|
||||
log "Running transformer tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
|
||||
;;
|
||||
"ssim")
|
||||
log "Running SSIM tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
;;
|
||||
"training")
|
||||
log "Running training tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests"
|
||||
;;
|
||||
"training_lora")
|
||||
log "Running LoRA training tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_lora_tests"
|
||||
;;
|
||||
"training_vsa")
|
||||
log "Running training VSA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests_VSA"
|
||||
;;
|
||||
"inference_sta")
|
||||
log "Running inference STA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_STA"
|
||||
;;
|
||||
"precision_sta")
|
||||
log "Running precision STA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_STA"
|
||||
;;
|
||||
"precision_vsa")
|
||||
log "Running precision VSA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_VSA"
|
||||
;;
|
||||
"inference_lora")
|
||||
log "Running LoRA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_lora_tests"
|
||||
;;
|
||||
"distillation_dmd")
|
||||
log "Running distillation DMD tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_distill_dmd_tests"
|
||||
;;
|
||||
# run_inference_tests_vmoba
|
||||
"inference_vmoba")
|
||||
log "Running V-MoBA inference tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_vmoba"
|
||||
;;
|
||||
"precision_vmoba")
|
||||
log "Running V-MoBA precision tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_vmoba"
|
||||
;;
|
||||
*)
|
||||
log "Error: Unknown test type: $TEST_TYPE"
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
|
||||
log "Executing: $MODAL_COMMAND"
|
||||
eval "$MODAL_COMMAND"
|
||||
TEST_EXIT_CODE=$?
|
||||
|
||||
if [ $TEST_EXIT_CODE -eq 0 ]; then
|
||||
log "Modal test completed successfully"
|
||||
else
|
||||
log "Error: Modal test failed with exit code: $TEST_EXIT_CODE"
|
||||
fi
|
||||
|
||||
log "=== Test execution completed with exit code: $TEST_EXIT_CODE ==="
|
||||
exit $TEST_EXIT_CODE
|
||||
@@ -1,40 +0,0 @@
|
||||
#!/bin/bash
|
||||
set -uo pipefail
|
||||
|
||||
log() {
|
||||
echo "[$(date '+%Y-%m-%d %H:%M:%S')] $1"
|
||||
}
|
||||
|
||||
log "=== Starting pre-commit checks ==="
|
||||
|
||||
cd "$(dirname "$0")/../.."
|
||||
PROJECT_ROOT=$(pwd)
|
||||
log "Project root: $PROJECT_ROOT"
|
||||
|
||||
if ! python3 -m pre_commit --version &> /dev/null; then
|
||||
log "pre-commit not found, installing..."
|
||||
python3 -m pip install --user pre-commit==4.0.1
|
||||
|
||||
if ! python3 -m pre_commit --version &> /dev/null; then
|
||||
log "Error: Failed to install pre-commit."
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
|
||||
log "Pre-commit version: $(python3 -m pre_commit --version)"
|
||||
|
||||
log "Installing/updating pre-commit hooks..."
|
||||
python3 -m pre_commit install --install-hooks
|
||||
|
||||
log "Running pre-commit checks on all files..."
|
||||
python3 -m pre_commit run --all-files
|
||||
PRE_COMMIT_EXIT_CODE=$?
|
||||
|
||||
if [ $PRE_COMMIT_EXIT_CODE -eq 0 ]; then
|
||||
log "Pre-commit checks completed successfully"
|
||||
else
|
||||
log "Error: Pre-commit checks failed with exit code: $PRE_COMMIT_EXIT_CODE"
|
||||
fi
|
||||
|
||||
log "=== Pre-commit checks completed with exit code: $PRE_COMMIT_EXIT_CODE ==="
|
||||
exit $PRE_COMMIT_EXIT_CODE
|
||||
@@ -4,6 +4,14 @@ title: "[Bug] "
|
||||
labels: ['Bug']
|
||||
|
||||
body:
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Environment
|
||||
description: |
|
||||
Please share your environment with us. You can run the command **python fastvideo/utils/env_utils.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Describe the bug
|
||||
@@ -17,13 +25,5 @@ body:
|
||||
What command or script did you run? Which **model** are you using?
|
||||
placeholder: |
|
||||
A placeholder for the command.
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Environment
|
||||
description: |
|
||||
Please share your environment with us. You can run the command **python collect_env.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
@@ -1,56 +0,0 @@
|
||||
name: 💬 Request for comments (RFC).
|
||||
description: Ask for feedback on major architectural changes or design choices.
|
||||
title: "[RFC]: "
|
||||
labels: ["RFC"]
|
||||
|
||||
body:
|
||||
- type: markdown
|
||||
attributes:
|
||||
value: >
|
||||
#### Please take a look at previous [RFCs](https://github.com/hao-ai-lab/FastVideo/issues?q=label%3ARFC+sort%3Aupdated-desc) for reference.
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Motivation.
|
||||
description: >
|
||||
The motivation of the RFC.
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Proposed Change.
|
||||
description: >
|
||||
The proposed change of the RFC.
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Feedback Period.
|
||||
description: >
|
||||
The feedback period of the RFC. Usually at least one week.
|
||||
validations:
|
||||
required: false
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: CC List.
|
||||
description: >
|
||||
The list of people you want to CC.
|
||||
validations:
|
||||
required: false
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Any Other Things.
|
||||
description: >
|
||||
Any other things you would like to mention.
|
||||
validations:
|
||||
required: false
|
||||
- type: markdown
|
||||
attributes:
|
||||
value: >
|
||||
Thanks for contributing 🎉!
|
||||
- type: checkboxes
|
||||
id: askllm
|
||||
attributes:
|
||||
label: Before submitting a new issue...
|
||||
options:
|
||||
- label: Make sure you already searched for relevant issues.
|
||||
required: true
|
||||
@@ -1,249 +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,
|
||||
required=True,
|
||||
help='Docker image to use')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
args = parse_arguments()
|
||||
API_KEY = os.environ['RUNPOD_API_KEY']
|
||||
RUN_ID = os.environ['GITHUB_RUN_ID']
|
||||
JOB_ID = os.environ['JOB_ID']
|
||||
PODS_API = "https://rest.runpod.io/v1/pods"
|
||||
HEADERS = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {API_KEY}"
|
||||
}
|
||||
|
||||
|
||||
def create_pod():
|
||||
"""Create a RunPod instance"""
|
||||
# Ensure image name is lowercase (Docker requirement)
|
||||
image_name = args.image.lower()
|
||||
print(f"Using specified image: {image_name}")
|
||||
|
||||
docker_start_cmd = [
|
||||
"bash",
|
||||
"-c",
|
||||
"apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
|
||||
]
|
||||
|
||||
print(f"Creating RunPod instance with GPU: {args.gpu_type}...")
|
||||
payload = {
|
||||
"name": f"fastvideo-{JOB_ID}-{RUN_ID}",
|
||||
"containerDiskInGb": args.disk_size,
|
||||
"volumeInGb": args.volume_size,
|
||||
"gpuTypeIds": [args.gpu_type],
|
||||
"gpuCount": args.gpu_count,
|
||||
"imageName": image_name,
|
||||
"allowedCudaVersions": ["12.4"],
|
||||
"dockerStartCmd": docker_start_cmd
|
||||
}
|
||||
|
||||
response = requests.post(PODS_API, headers=HEADERS, json=payload)
|
||||
response_data = response.json()
|
||||
print(f"Response: {json.dumps(response_data, indent=2)}")
|
||||
|
||||
return response_data["id"]
|
||||
|
||||
|
||||
def wait_for_pod(pod_id):
|
||||
"""Wait for pod to be in RUNNING state and fully ready with SSH access"""
|
||||
print("Waiting for RunPod to be ready...")
|
||||
|
||||
# First wait for RUNNING status
|
||||
max_attempts = 10
|
||||
attempts = 0
|
||||
while attempts < max_attempts:
|
||||
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
pod_data = response.json()
|
||||
status = pod_data["desiredStatus"]
|
||||
|
||||
if status == "RUNNING":
|
||||
print("RunPod is running! Now waiting for ports to be assigned...")
|
||||
break
|
||||
|
||||
print(
|
||||
f"Current status: {status}, waiting... (attempt {attempts+1}/{max_attempts})"
|
||||
)
|
||||
time.sleep(2)
|
||||
attempts += 1
|
||||
|
||||
if attempts >= max_attempts:
|
||||
raise TimeoutError(
|
||||
"Timed out waiting for RunPod to reach RUNNING state")
|
||||
|
||||
# Wait for ports to be assigned
|
||||
max_attempts = 50
|
||||
attempts = 0
|
||||
while attempts < max_attempts:
|
||||
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
pod_data = response.json()
|
||||
port_mappings = pod_data.get("portMappings")
|
||||
|
||||
if (port_mappings is not None and "22" in port_mappings
|
||||
and pod_data.get("publicIp", "") != ""):
|
||||
print("RunPod is ready with SSH access!")
|
||||
print(f"SSH IP: {pod_data['publicIp']}")
|
||||
print(f"SSH Port: {port_mappings['22']}")
|
||||
break
|
||||
|
||||
print(
|
||||
f"Waiting for SSH port and public IP to be available... (attempt {attempts+1}/{max_attempts})"
|
||||
)
|
||||
time.sleep(20)
|
||||
attempts += 1
|
||||
|
||||
if attempts >= max_attempts:
|
||||
raise TimeoutError("Timed out waiting for RunPod SSH access")
|
||||
|
||||
|
||||
def execute_command(pod_id):
|
||||
"""Execute command on the pod via SSH using system SSH client"""
|
||||
print(f"Running command: {args.test_command}")
|
||||
|
||||
response = requests.get(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
pod_data = response.json()
|
||||
ssh_ip = pod_data["publicIp"]
|
||||
ssh_port = pod_data["portMappings"]["22"]
|
||||
|
||||
# Copy the repository to the pod using scp
|
||||
repo_dir = os.path.abspath(os.getcwd())
|
||||
repo_name = os.path.basename(repo_dir)
|
||||
|
||||
print(f"Copying repository from {repo_dir} to RunPod...")
|
||||
|
||||
tar_command = [
|
||||
"tar", "-czf", "/tmp/repo.tar.gz", "-C",
|
||||
os.path.dirname(repo_dir), repo_name
|
||||
]
|
||||
subprocess.run(tar_command, check=True)
|
||||
|
||||
# Copy the tarball to the pod
|
||||
scp_command = [
|
||||
"scp", "-o", "StrictHostKeyChecking=no", "-o",
|
||||
"UserKnownHostsFile=/dev/null", "-o", "ServerAliveInterval=60", "-o",
|
||||
"ServerAliveCountMax=10", "-P",
|
||||
str(ssh_port), "/tmp/repo.tar.gz", f"root@{ssh_ip}:/tmp/"
|
||||
]
|
||||
subprocess.run(scp_command, check=True)
|
||||
|
||||
# For custom image, we can use the pre-configured environment
|
||||
setup_steps = [
|
||||
"tar -xzf /tmp/repo.tar.gz --no-same-owner -C /workspace/",
|
||||
f"cd /workspace/{repo_name}",
|
||||
"source $HOME/.local/bin/env && source /opt/venv/bin/activate",
|
||||
args.test_command
|
||||
]
|
||||
|
||||
remote_command = " && ".join(setup_steps)
|
||||
|
||||
ssh_command = [
|
||||
"ssh", "-o", "StrictHostKeyChecking=no", "-o",
|
||||
"UserKnownHostsFile=/dev/null", "-o", "ServerAliveInterval=60", "-o",
|
||||
"ServerAliveCountMax=10", "-p",
|
||||
str(ssh_port), f"root@{ssh_ip}", remote_command
|
||||
]
|
||||
|
||||
print(f"Connecting to {ssh_ip}:{ssh_port}...")
|
||||
|
||||
try:
|
||||
process = subprocess.Popen(ssh_command,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=0)
|
||||
|
||||
stdout_lines = []
|
||||
|
||||
print("Command output:")
|
||||
|
||||
for line in iter(process.stdout.readline, ''):
|
||||
print(line.strip())
|
||||
stdout_lines.append(line)
|
||||
|
||||
process.wait()
|
||||
|
||||
return_code = process.returncode
|
||||
success = return_code == 0
|
||||
|
||||
stdout_str = "".join(stdout_lines)
|
||||
|
||||
if success:
|
||||
print("Command executed successfully")
|
||||
else:
|
||||
print(f"Command failed with exit code {return_code}")
|
||||
|
||||
result = {
|
||||
"success": success,
|
||||
"return_code": return_code,
|
||||
"stdout": stdout_str,
|
||||
"stderr": ""
|
||||
}
|
||||
return result
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error executing SSH command: {str(e)}")
|
||||
result = {"success": False, "error": str(e), "stdout": "", "stderr": ""}
|
||||
return result
|
||||
|
||||
|
||||
def terminate_pod(pod_id):
|
||||
"""Terminate the pod"""
|
||||
print("Terminating RunPod...")
|
||||
requests.delete(f"{PODS_API}/{pod_id}", headers=HEADERS)
|
||||
print(f"Terminated pod {pod_id}")
|
||||
|
||||
|
||||
def main():
|
||||
pod_id = None
|
||||
try:
|
||||
pod_id = create_pod()
|
||||
wait_for_pod(pod_id)
|
||||
result = execute_command(pod_id)
|
||||
|
||||
if result.get("error") is not None:
|
||||
print(f"Error executing command: {result['error']}")
|
||||
sys.exit(1)
|
||||
|
||||
if not result.get("success", False):
|
||||
print(
|
||||
"Tests failed - check the output above for details on which tests failed"
|
||||
)
|
||||
sys.exit(1)
|
||||
|
||||
finally:
|
||||
if pod_id:
|
||||
terminate_pod(pod_id)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,90 +0,0 @@
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import uuid
|
||||
|
||||
import requests
|
||||
|
||||
API_KEY = os.environ['RUNPOD_API_KEY']
|
||||
RUN_ID = os.environ.get('GITHUB_RUN_ID', str(uuid.uuid4()))
|
||||
PODS_API = "https://rest.runpod.io/v1/pods"
|
||||
HEADERS = {
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": f"Bearer {API_KEY}"
|
||||
}
|
||||
|
||||
|
||||
def get_job_ids():
|
||||
"""Parse job IDs from environment variable"""
|
||||
job_ids_str = os.environ.get('JOB_IDS')
|
||||
try:
|
||||
job_ids = json.loads(job_ids_str)
|
||||
if not isinstance(job_ids, list):
|
||||
print("Error: JOB_IDS is not a list.")
|
||||
sys.exit(1)
|
||||
return job_ids
|
||||
except json.JSONDecodeError as e:
|
||||
print(f"Error parsing JOB_IDS: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def cleanup_pods():
|
||||
"""Find and terminate RunPod instances"""
|
||||
print(f"Run ID: {RUN_ID}")
|
||||
|
||||
single_job_id = os.environ.get('JOB_ID')
|
||||
|
||||
if single_job_id:
|
||||
job_ids = [single_job_id]
|
||||
print(f"Job ID: {single_job_id}")
|
||||
else:
|
||||
job_ids = get_job_ids()
|
||||
print(f"Job IDs: {job_ids}")
|
||||
|
||||
# Get all pods associated with RunPod API_KEY
|
||||
try:
|
||||
response = requests.get(PODS_API, headers=HEADERS)
|
||||
response.raise_for_status()
|
||||
pods = response.json()
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"Error getting pods: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
# Find and terminate pods created by this workflow run
|
||||
terminated_pods = []
|
||||
for pod in pods:
|
||||
pod_name = pod.get("name", "")
|
||||
pod_id = pod.get("id")
|
||||
|
||||
# Check if this pod was created by one of our jobs
|
||||
if any(f"{job_id}-{RUN_ID}" in pod_name for job_id in job_ids):
|
||||
print(f"Found pod: {pod_id} ({pod_name})")
|
||||
try:
|
||||
print(f"Terminating pod {pod_id}...")
|
||||
term_response = requests.delete(f"{PODS_API}/{pod_id}",
|
||||
headers=HEADERS)
|
||||
term_response.raise_for_status()
|
||||
terminated_pods.append(pod_id)
|
||||
print(f"Successfully terminated pod {pod_id}")
|
||||
except requests.exceptions.RequestException as e:
|
||||
print(f"Error terminating pod {pod_id}: {e}")
|
||||
sys.exit(1)
|
||||
|
||||
if terminated_pods:
|
||||
if single_job_id:
|
||||
print(f"Terminated pod: {terminated_pods[0]}")
|
||||
else:
|
||||
print(f"Terminated {len(terminated_pods)} pods: {terminated_pods}")
|
||||
else:
|
||||
if single_job_id:
|
||||
print(f"No pod found matching pattern: {single_job_id}-{RUN_ID}")
|
||||
else:
|
||||
print("No pods found to terminate.")
|
||||
|
||||
|
||||
def main():
|
||||
cleanup_pods()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,106 +0,0 @@
|
||||
name: Build Image Template
|
||||
|
||||
on:
|
||||
workflow_call:
|
||||
inputs:
|
||||
python_version:
|
||||
required: true
|
||||
type: string
|
||||
dockerfile_path:
|
||||
required: true
|
||||
type: string
|
||||
tag_suffix:
|
||||
required: true
|
||||
type: string
|
||||
|
||||
jobs:
|
||||
build-and-push:
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
# Display initial space
|
||||
echo "Initial disk space:"
|
||||
df -h
|
||||
|
||||
# Remove large directories directly
|
||||
sudo rm -rf /usr/share/dotnet
|
||||
sudo rm -rf /usr/local/lib/android
|
||||
sudo rm -rf /opt/ghc
|
||||
sudo rm -rf /usr/local/share/boost
|
||||
sudo rm -rf /usr/share/swift
|
||||
sudo rm -rf /usr/local/lib/node_modules
|
||||
sudo rm -rf /usr/local/share/powershell
|
||||
sudo rm -rf /usr/share/rust
|
||||
sudo rm -rf /usr/local/.ghcup
|
||||
|
||||
# Remove cached files
|
||||
sudo rm -rf /var/lib/apt/lists/*
|
||||
sudo rm -rf /var/cache/apt/archives/*
|
||||
|
||||
# Clean Docker
|
||||
docker system prune -af --volumes
|
||||
|
||||
# Display available space after cleanup
|
||||
echo "Disk space after cleanup:"
|
||||
df -h
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@v3
|
||||
|
||||
- name: Login to GitHub Container Registry
|
||||
uses: docker/login-action@v3
|
||||
with:
|
||||
registry: ghcr.io
|
||||
username: ${{ github.repository_owner }}
|
||||
password: ${{ secrets.GITHUB_TOKEN }}
|
||||
|
||||
- name: Prepare tags
|
||||
id: prepare-tags
|
||||
run: |
|
||||
SHORT_SHA=$(echo ${{ github.sha }} | cut -c1-7)
|
||||
|
||||
TAGS="type=raw,value=${{ inputs.tag_suffix }}-latest"
|
||||
TAGS="${TAGS}\ntype=raw,value=${{ inputs.tag_suffix }}-sha-${SHORT_SHA}"
|
||||
|
||||
# Set Python 3.10 as the default image
|
||||
if [[ "${{ inputs.python_version }}" == "3.10" ]]; then
|
||||
TAGS="${TAGS}\ntype=raw,value=latest"
|
||||
fi
|
||||
|
||||
{
|
||||
echo "tags<<EOF"
|
||||
echo -e "$TAGS"
|
||||
echo "EOF"
|
||||
} >> $GITHUB_OUTPUT
|
||||
|
||||
- name: Extract metadata for Docker
|
||||
id: meta
|
||||
uses: docker/metadata-action@v5
|
||||
with:
|
||||
images: ghcr.io/${{ github.repository }}/fastvideo-dev
|
||||
tags: ${{ steps.prepare-tags.outputs.tags }}
|
||||
|
||||
- name: Build and push Docker image
|
||||
id: build-push
|
||||
uses: docker/build-push-action@v6
|
||||
with:
|
||||
context: .
|
||||
file: ${{ inputs.dockerfile_path }}
|
||||
push: true
|
||||
tags: ${{ steps.meta.outputs.tags }}
|
||||
labels: ${{ steps.meta.outputs.labels }}
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
|
||||
- name: Success message
|
||||
run: |
|
||||
echo "✅ Python ${{ inputs.python_version }} image successfully built and pushed to ghcr.io/${{ github.repository }}/fastvideo-dev:${{ inputs.tag_suffix }}-latest"
|
||||
echo "To run tests with this image, manually trigger the 'Run Tests' workflow."
|
||||
@@ -1,67 +0,0 @@
|
||||
name: Build and Push Docker Images
|
||||
|
||||
on:
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
python_3_10:
|
||||
description: 'Build Python 3.10 image'
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
python_3_11:
|
||||
description: 'Build Python 3.11 image'
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
python_3_12:
|
||||
description: 'Build Python 3.12 image'
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
python_3_12_cuda_12_9:
|
||||
description: 'Build Python 3.12 image Cuda 12.9'
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
packages: write
|
||||
|
||||
jobs:
|
||||
build-python-3-10:
|
||||
if: ${{ github.event.inputs.python_3_10 == 'true' }}
|
||||
uses: ./.github/workflows/build-image-template.yml
|
||||
with:
|
||||
python_version: '3.10'
|
||||
dockerfile_path: docker/Dockerfile.python3.10
|
||||
tag_suffix: py3.10
|
||||
secrets: inherit
|
||||
|
||||
build-python-3-11:
|
||||
if: ${{ github.event.inputs.python_3_11 == 'true' }}
|
||||
uses: ./.github/workflows/build-image-template.yml
|
||||
with:
|
||||
python_version: '3.11'
|
||||
dockerfile_path: docker/Dockerfile.python3.11
|
||||
tag_suffix: py3.11
|
||||
secrets: inherit
|
||||
|
||||
build-python-3-12:
|
||||
if: ${{ github.event.inputs.python_3_12 == 'true' }}
|
||||
uses: ./.github/workflows/build-image-template.yml
|
||||
with:
|
||||
python_version: '3.12'
|
||||
dockerfile_path: docker/Dockerfile.python3.12
|
||||
tag_suffix: py3.12
|
||||
secrets: inherit
|
||||
|
||||
build-python-3-12-cuda-12-9:
|
||||
if: ${{ github.event.inputs.python_3_12_cuda_12_9 == 'true' }}
|
||||
uses: ./.github/workflows/build-image-template.yml
|
||||
with:
|
||||
python_version: '3.12'
|
||||
dockerfile_path: docker/Dockerfile.python3.12.cuda12.9.1
|
||||
tag_suffix: py3.12-cuda12.9.1
|
||||
secrets: inherit
|
||||
@@ -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,83 +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"
|
||||
- "fastvideo/examples/**/*.py"
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
types: [opened, ready_for_review, synchronize, reopened]
|
||||
paths:
|
||||
- "docs/**/*.md"
|
||||
- "fastvideo/examples/**/*.py"
|
||||
|
||||
# Allows you to run this workflow manually from the Actions tab
|
||||
workflow_dispatch:
|
||||
|
||||
# Sets permissions of the GITHUB_TOKEN to allow deployment to GitHub Pages
|
||||
permissions:
|
||||
contents: read
|
||||
pages: write
|
||||
id-token: write
|
||||
|
||||
# Allow only one concurrent deployment, skipping runs queued between the run in-progress and latest queued.
|
||||
# However, do NOT cancel in-progress runs as we want to allow these production deployments to complete.
|
||||
concurrency:
|
||||
group: "pages"
|
||||
cancel-in-progress: false
|
||||
|
||||
# Default to bash
|
||||
defaults:
|
||||
run:
|
||||
shell: bash
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
uses: ./.github/workflows/pre-commit.yml
|
||||
|
||||
# Build job
|
||||
build:
|
||||
runs-on: ubuntu-latest
|
||||
needs: pre-commit
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
- name: Setup Pages
|
||||
id: pages
|
||||
uses: actions/configure-pages@v5
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
cd docs
|
||||
pip install -r requirements-docs.txt
|
||||
- name: Build docs
|
||||
run: |
|
||||
cd docs
|
||||
make clean
|
||||
make html
|
||||
- name: Upload artifact
|
||||
uses: actions/upload-pages-artifact@v3
|
||||
with:
|
||||
path: ./docs/build/html
|
||||
|
||||
# Deployment job
|
||||
deploy:
|
||||
environment:
|
||||
name: github-pages
|
||||
url: ${{ steps.deployment.outputs.page_url }}
|
||||
if: ${{ github.event_name == 'push' }}
|
||||
runs-on: ubuntu-latest
|
||||
needs: build
|
||||
steps:
|
||||
- name: Deploy to GitHub Pages
|
||||
id: deployment
|
||||
uses: actions/deploy-pages@v4
|
||||
@@ -6,7 +6,6 @@ on:
|
||||
- main
|
||||
paths:
|
||||
- 'pyproject.toml' # Trigger when pyproject.toml changes
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
check-version-change:
|
||||
@@ -16,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
|
||||
|
||||
@@ -24,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
|
||||
@@ -42,17 +41,17 @@ jobs:
|
||||
|
||||
build-publish-main:
|
||||
needs: check-version-change
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
if: needs.check-version-change.outputs.version-changed == 'true'
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
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,376 +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"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
- "csrc/**"
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
run_encoder_test:
|
||||
description: "Run encoder-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_vae_test:
|
||||
description: "Run vae-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_transformer_test:
|
||||
description: "Run transformer-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_ssim_test:
|
||||
description: "Run ssim-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_training_test:
|
||||
description: "Run training-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_training_test_VSA:
|
||||
description: "Run training-test-VSA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_inference_test_STA:
|
||||
description: "Run inference-test-STA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_precision_test_STA:
|
||||
description: "Run precision-test-STA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_precision_test_VSA:
|
||||
description: "Run precision-test-VSA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_nightly_test:
|
||||
description: "Run nightly-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
|
||||
env:
|
||||
PYTHONUNBUFFERED: "1"
|
||||
|
||||
|
||||
concurrency:
|
||||
group: pr-test-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
pre-commit:
|
||||
uses: ./.github/workflows/pre-commit.yml
|
||||
|
||||
change-filter:
|
||||
runs-on: ubuntu-latest
|
||||
needs: pre-commit
|
||||
if: ${{ github.event.pull_request.draft == false || github.event_name == 'workflow_dispatch' }}
|
||||
outputs:
|
||||
encoder-test: ${{ steps.filter.outputs.encoder-test }}
|
||||
vae-test: ${{ steps.filter.outputs.vae-test }}
|
||||
transformer-test: ${{ steps.filter.outputs.transformer-test }}
|
||||
training-test: ${{ steps.filter.outputs.training-test }}
|
||||
training-test-VSA: ${{ steps.filter.outputs.training-test-VSA }}
|
||||
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
|
||||
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
|
||||
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
id: filter
|
||||
with:
|
||||
filters: |
|
||||
# Define reusable path patterns
|
||||
common-paths: &common-paths
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
sta-kernel-paths: &sta-kernel-paths
|
||||
- 'csrc/attn/sliding_tile_attn/**'
|
||||
- 'csrc/attn/sliding_tile_attn/tk/**'
|
||||
- 'csrc/attn/sliding_tile_attn/setup.py'
|
||||
- 'csrc/attn/sliding_tile_attn/config_sta.py'
|
||||
- 'csrc/attn/sliding_tile_attn/st_attn.cpp'
|
||||
vsa-kernel-paths: &vsa-kernel-paths
|
||||
- 'csrc/attn/video_sparse_attn/**'
|
||||
- 'csrc/attn/video_sparse_attn/tk/**'
|
||||
- 'csrc/attn/video_sparse_attn/setup.py'
|
||||
- 'csrc/attn/video_sparse_attn/config_vsa.py'
|
||||
- 'csrc/attn/video_sparse_attn/vsa.cpp'
|
||||
vsa-paths: &vsa-paths
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
|
||||
# Actual tests
|
||||
encoder-test:
|
||||
- 'fastvideo/models/encoders/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/encoders/**'
|
||||
- *common-paths
|
||||
vae-test:
|
||||
- 'fastvideo/models/vaes/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/vaes/**'
|
||||
- *common-paths
|
||||
transformer-test:
|
||||
- 'fastvideo/models/dits/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/transformers/**'
|
||||
- 'fastvideo/layers/**'
|
||||
- 'fastvideo/attention/**'
|
||||
- *common-paths
|
||||
training-test:
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
training-test-VSA:
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
inference-test-STA:
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
- *sta-kernel-paths
|
||||
precision-test-STA:
|
||||
- *common-paths
|
||||
- *sta-kernel-paths
|
||||
precision-test-VSA:
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
|
||||
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')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "encoder-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/encoders -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
vae-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.vae-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_vae_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "vae-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/vaes -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
transformer-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.transformer-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_transformer_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "transformer-test"
|
||||
gpu_type: "NVIDIA L40S"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/transformers -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
ssim-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
github.event_name != 'workflow_dispatch' || (github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
python-version: [
|
||||
# {version: "3.10", tag: "latest"},
|
||||
# {version: "3.11", tag: "py3.11-latest"},
|
||||
{version: "3.12", tag: "py3.12-latest"}
|
||||
]
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "ssim-test-py${{ matrix.python-version.version }}"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 2
|
||||
volume_size: 200
|
||||
disk_size: 200
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:${{ matrix.python-version.tag }}"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/ssim -vs"
|
||||
timeout_minutes: 60
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
training-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "training-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/training/Vanilla -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
training-test-VSA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test-VSA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test_VSA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "training-test-VSA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 2
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/training/VSA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
inference-test-STA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.inference-test-STA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_inference_test_STA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "inference-test-STA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 2
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/inference/STA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
precision-test-STA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-STA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_STA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "precision-test-STA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_sta.py"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
precision-test-VSA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-VSA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_VSA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "precision-test-VSA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_vsa.py"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
nightly-test:
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "nightly-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
runpod-cleanup:
|
||||
# Add other jobs to this list as you create them
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test, training-test, training-test-VSA, inference-test-STA, precision-test-STA, precision-test-VSA]
|
||||
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
|
||||
- name: Install dependencies
|
||||
run: pip install requests
|
||||
|
||||
- name: Cleanup all RunPod instances
|
||||
env:
|
||||
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12", "training-test", "training-test-VSA", "inference-test-STA", "precision-test-STA", "precision-test-VSA"]'
|
||||
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.12"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/actionlint.json"
|
||||
- run: echo "::add-matcher::.github/workflows/matchers/mypy.json"
|
||||
- uses: pre-commit/action@v3.0.1
|
||||
with:
|
||||
extra_args: --all-files --hook-stage manual
|
||||
@@ -1,28 +0,0 @@
|
||||
name: Publish to Comfy registry
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- master
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ github.repository_owner == 'hao-ai-lab' }}
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
submodules: true
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
@@ -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
|
||||
@@ -1,94 +0,0 @@
|
||||
name: RunPod Test
|
||||
|
||||
on:
|
||||
workflow_call:
|
||||
inputs:
|
||||
job_id:
|
||||
required: true
|
||||
type: string
|
||||
description: "Unique identifier for this test job"
|
||||
gpu_type:
|
||||
required: true
|
||||
type: string
|
||||
description: "GPU type to use (e.g. NVIDIA A40, NVIDIA L40S)"
|
||||
gpu_count:
|
||||
required: true
|
||||
type: number
|
||||
description: "Number of GPUs to use"
|
||||
volume_size:
|
||||
required: false
|
||||
type: number
|
||||
default: 20
|
||||
description: "Volume size in GB"
|
||||
disk_size:
|
||||
required: false
|
||||
type: number
|
||||
default: 20
|
||||
description: "Disk size in GB"
|
||||
image:
|
||||
required: true
|
||||
type: string
|
||||
description: "Docker image to use"
|
||||
test_command:
|
||||
required: true
|
||||
type: string
|
||||
description: "Command to run tests"
|
||||
timeout_minutes:
|
||||
required: false
|
||||
type: number
|
||||
default: 30
|
||||
description: "Timeout in minutes"
|
||||
secrets:
|
||||
RUNPOD_API_KEY:
|
||||
required: true
|
||||
RUNPOD_PRIVATE_KEY:
|
||||
required: true
|
||||
WANDB_API_KEY:
|
||||
required: false
|
||||
|
||||
jobs:
|
||||
run-test:
|
||||
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.12"
|
||||
|
||||
- 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: ${{ inputs.job_id }}
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
timeout-minutes: ${{ inputs.timeout_minutes }}
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
--gpu-type "${{ inputs.gpu_type }}"
|
||||
--gpu-count ${{ inputs.gpu_count }}
|
||||
--volume-size ${{ inputs.volume_size }}
|
||||
--disk-size ${{ inputs.disk_size }}
|
||||
--image "${{ inputs.image }}"
|
||||
--test-command "${{ inputs.test_command }}"
|
||||
|
||||
- name: Terminate RunPod Instances
|
||||
if: ${{ always() }}
|
||||
env:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
JOB_ID: ${{ inputs.job_id }}
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
@@ -5,8 +5,7 @@ on:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
workflow_dispatch:
|
||||
- "csrc/sliding_tile_attention/setup.py"
|
||||
|
||||
jobs:
|
||||
check-version-change:
|
||||
@@ -16,14 +15,14 @@ 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
|
||||
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/attn/sliding_tile_attn
|
||||
cd csrc/sliding_tile_attention
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
echo "New version: $NEW_VERSION"
|
||||
@@ -44,7 +43,7 @@ jobs:
|
||||
build_wheels:
|
||||
name: Build Wheel
|
||||
needs: check-version-change
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
if: needs.check-version-change.outputs.version-changed == 'true'
|
||||
runs-on: ${{ matrix.os }}
|
||||
|
||||
strategy:
|
||||
@@ -58,29 +57,6 @@ jobs:
|
||||
cuda-version: ['12.4.1', '12.5.1', '12.6.3']
|
||||
|
||||
steps:
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
echo "Initial disk space:"
|
||||
df -h
|
||||
|
||||
# Remove large directories
|
||||
sudo rm -rf /usr/share/dotnet
|
||||
sudo rm -rf /usr/local/lib/android
|
||||
sudo rm -rf /opt/ghc
|
||||
sudo rm -rf /usr/local/share/boost
|
||||
sudo rm -rf /usr/share/swift
|
||||
sudo rm -rf /usr/local/lib/node_modules
|
||||
sudo rm -rf /usr/local/share/powershell
|
||||
sudo rm -rf /usr/share/rust
|
||||
sudo rm -rf /usr/local/.ghcup
|
||||
|
||||
# Remove cached files
|
||||
sudo rm -rf /var/lib/apt/lists/*
|
||||
sudo rm -rf /var/cache/apt/archives/*
|
||||
|
||||
echo "Disk space after cleanup:"
|
||||
df -h
|
||||
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
|
||||
@@ -136,21 +112,19 @@ jobs:
|
||||
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn/sliding_tile_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
cd csrc/sliding_tile_attention # Move into the correct folder
|
||||
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/attn/sliding_tile_attn
|
||||
cd csrc/sliding_tile_attention
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-version }} | cut -d. -f1,2)
|
||||
@@ -165,13 +139,13 @@ jobs:
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}
|
||||
path: csrc/attn/sliding_tile_attn/dist/*.whl
|
||||
path: csrc/sliding_tile_attention/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
name: Publish package
|
||||
needs: [build_wheels, check-version-change]
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
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
|
||||
@@ -231,19 +205,17 @@ jobs:
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn/sliding_tile_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
cd csrc/sliding_tile_attention # Move into the correct folder
|
||||
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/attn/sliding_tile_attn/dist/
|
||||
packages-dir: csrc/sliding_tile_attention/dist/
|
||||
|
||||
@@ -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/attn/test
|
||||
pytest --ignore csrc/sliding_tile_attention/test
|
||||
@@ -1,257 +0,0 @@
|
||||
name: Publish Video Sparse Attention Kernel to PyPI on Version Change
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
check-version-change:
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
version-changed: ${{ steps.check-version.outputs.changed }}
|
||||
new-version: ${{ steps.check-version.outputs.new-version }}
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 2
|
||||
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/attn/video_sparse_attn
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
|
||||
echo "changed=true" >> $GITHUB_OUTPUT
|
||||
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
|
||||
else
|
||||
echo "Version did not change"
|
||||
echo "changed=false" >> $GITHUB_OUTPUT
|
||||
fi
|
||||
|
||||
build_wheels:
|
||||
name: Build Wheel
|
||||
needs: check-version-change
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ${{ matrix.os }}
|
||||
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
# Using ubuntu-20.04 instead of 22.04 for more compatibility (glibc). Ideally we'd use the
|
||||
# manylinux docker image, but I haven't figured out how to install CUDA on manylinux.
|
||||
os: [ubuntu-22.04]
|
||||
python-version: ['3.10', '3.11', '3.12', '3.13']
|
||||
# For version reference https://pytorch.org/get-started/previous-versions/
|
||||
torch-cuda:
|
||||
- torch-version: '2.5.1'
|
||||
cuda-version: '12.4.1'
|
||||
torch-cuda-short: 'cu124'
|
||||
- torch-version: '2.6.0'
|
||||
cuda-version: '12.6.3'
|
||||
torch-cuda-short: 'cu126'
|
||||
- torch-version: '2.7.1'
|
||||
cuda-version: '12.8.0'
|
||||
torch-cuda-short: 'cu128'
|
||||
|
||||
steps:
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
echo "Initial disk space:"
|
||||
df -h
|
||||
|
||||
# Remove large directories
|
||||
sudo rm -rf /usr/share/dotnet
|
||||
sudo rm -rf /usr/local/lib/android
|
||||
sudo rm -rf /opt/ghc
|
||||
sudo rm -rf /usr/local/share/boost
|
||||
sudo rm -rf /usr/share/swift
|
||||
sudo rm -rf /usr/local/lib/node_modules
|
||||
sudo rm -rf /usr/local/share/powershell
|
||||
sudo rm -rf /usr/share/rust
|
||||
sudo rm -rf /usr/local/.ghcup
|
||||
|
||||
# Remove cached files
|
||||
sudo rm -rf /var/lib/apt/lists/*
|
||||
sudo rm -rf /var/cache/apt/archives/*
|
||||
|
||||
echo "Disk space after cleanup:"
|
||||
df -h
|
||||
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Install CUDA ${{ matrix.torch-cuda.cuda-version }}
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: ${{ matrix.torch-cuda.cuda-version }}
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-${{ matrix.torch-cuda.cuda-version }}
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch ${{ matrix.torch-cuda.torch-version }}+cu${{ matrix.torch-cuda.cuda-version }}
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
|
||||
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
|
||||
pip install typing-extensions==4.12.2
|
||||
# We want to figure out the CUDA version to download pytorch
|
||||
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
|
||||
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
|
||||
pip install --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn/video_sparse_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/attn/video_sparse_attn
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.torch-version }} | cut -d. -f1,2)
|
||||
# Get the correct version format
|
||||
tmpname=cu${CUDA_SHORT_VERSION}torch${TORCH_SHORT_VERSION}
|
||||
wheel_name=$(ls dist/*whl | xargs -n 1 basename | sed "s/-/+$tmpname-/2")
|
||||
# Rename with version information
|
||||
ls dist/*whl |xargs -I {} mv {} dist/${wheel_name}
|
||||
echo "wheel_name=${wheel_name}" >> $GITHUB_ENV
|
||||
|
||||
- name: Upload wheel artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}
|
||||
path: csrc/attn/video_sparse_attn/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
name: Publish package
|
||||
needs: [build_wheels, check-version-change]
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ubuntu-22.04
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
- name: Install CUDA 12.4.1
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: 12.4.1
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
sub-packages: '["nvcc"]'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-12.4.1
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch 2.5.1+cu12.4.1
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
|
||||
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
|
||||
pip install typing-extensions==4.12.2
|
||||
# We want to figure out the CUDA version to download pytorch
|
||||
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
|
||||
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
|
||||
export TORCH_CUDA_VERSION=124
|
||||
pip install --no-cache-dir torch==2.5.1 --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn/video_sparse_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/attn/video_sparse_attn/dist/
|
||||
@@ -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,8 +1,11 @@
|
||||
__pycache__
|
||||
*.mp4
|
||||
.ipynb_checkpoints
|
||||
*.pth
|
||||
UCF-101/
|
||||
results/
|
||||
build/
|
||||
fastvideo.egg-info/
|
||||
wandb/
|
||||
*.ipynb
|
||||
*.jpg
|
||||
@@ -20,48 +23,14 @@ samples/
|
||||
data/
|
||||
outputs/
|
||||
outputs_video
|
||||
checkpoints/
|
||||
sbatch.sh
|
||||
*.out
|
||||
env
|
||||
dist/
|
||||
*.o
|
||||
**/build/
|
||||
**.egg-info
|
||||
**.pyc
|
||||
**.egg
|
||||
**.txt
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
dist/
|
||||
*.egg-info/
|
||||
*.egg
|
||||
eggs/
|
||||
.eggs/
|
||||
|
||||
# Sphinx documentation
|
||||
docs/_build/
|
||||
docs/source/getting_started/examples/
|
||||
docs/source/inference/examples/
|
||||
docs/source/training/examples/
|
||||
docs/source/distillation/examples/
|
||||
|
||||
# VSCode
|
||||
.vscode/
|
||||
|
||||
# DS Store
|
||||
.DS_Store
|
||||
|
||||
# vim swap files
|
||||
*.swo
|
||||
*.swp
|
||||
|
||||
# Python pickle files
|
||||
*.pkl
|
||||
|
||||
# Reference videos
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.mp4
|
||||
|
||||
# Static images
|
||||
!docs/source/_static/images/**/*.png
|
||||
!comfyui/assets/**/*.png
|
||||
!comfyui/assets/**/*.gif
|
||||
dmd_t2v_output/
|
||||
**.json
|
||||
@@ -1,7 +1,3 @@
|
||||
[submodule "csrc/attn/video_sparse_attn/tk"]
|
||||
path = csrc/attn/video_sparse_attn/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
[submodule "csrc/attn/sliding_tile_attn/tk"]
|
||||
path = csrc/attn/sliding_tile_attn/tk
|
||||
[submodule "csrc/sliding_tile_attention/tk"]
|
||||
path = sta_kernel/thunderkitten/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
@@ -1,84 +0,0 @@
|
||||
default_stages:
|
||||
- pre-commit # Run locally
|
||||
- manual # Run in CI
|
||||
exclude: |
|
||||
(?x)(
|
||||
fastvideo/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/.*|
|
||||
examples/.*|
|
||||
.github/workflows/fastvideo-publish.yml|
|
||||
.github/workflows/sta-publish.yml|
|
||||
.github/workflows/vsa-publish.yml|
|
||||
.github/workflows/build-image-template.yml|
|
||||
docs/source/inference/support_matrix.md
|
||||
)
|
||||
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.12
|
||||
hooks:
|
||||
- id: ruff
|
||||
args: [--output-format, github, --fix]
|
||||
- repo: https://github.com/codespell-project/codespell
|
||||
rev: v2.4.1
|
||||
hooks:
|
||||
- id: codespell
|
||||
additional_dependencies: ['tomli']
|
||||
args: ['--toml', 'pyproject.toml']
|
||||
- repo: https://github.com/PyCQA/isort
|
||||
rev: 6.0.1
|
||||
hooks:
|
||||
- id: isort
|
||||
- repo: https://github.com/jackdewinter/pymarkdown
|
||||
rev: v0.9.30
|
||||
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", "--disable-error-code", "union-attr", "--disable-error-code", "override" ]
|
||||
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/tests/ssim/" | grep -v "^fastvideo/tests/inference/lora/L40S_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,170 +1,230 @@
|
||||
<div align="center">
|
||||
<img src=assets/logos/logo.svg width="30%"/>
|
||||
<img src=assets/logo.jpg width="30%"/>
|
||||
</div>
|
||||
|
||||
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
|
||||
FastVideo is a lightweight framework for accelerating large video diffusion models.
|
||||
|
||||
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
|
||||
|
||||
<p align="center">
|
||||
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/S7HLCSTh" target="_blank"> <b> WeChat </b> </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>
|
||||
|
||||
<div align="center">
|
||||
<img src=assets/fastwan.png width="90%"/>
|
||||
</div>
|
||||
|
||||
## NEWS
|
||||
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
- ```2025/02/18```: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
|
||||
## Key Features
|
||||
|
||||
FastVideo has the following features:
|
||||
- End-to-end post-training support:
|
||||
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 to achineve >50x denoising speedup
|
||||
- Data preprocessing pipeline for video data
|
||||
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs
|
||||
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs
|
||||
- State-of-the-art performance optimizations for inference
|
||||
- [Video Sparse Attention](https://arxiv.org/pdf/2505.13389)
|
||||
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
|
||||
- [TeaCache](https://arxiv.org/pdf/2411.19108)
|
||||
- [Sage Attention](https://arxiv.org/abs/2410.02367)
|
||||
- Diverse hardware and OS support
|
||||
- Support H100, A100, 4090
|
||||
- Support Linux, Windows, MacOS
|
||||
|
||||
## Getting Started
|
||||
We recommend using an environment manager such as `Conda` to create a clean environment:
|
||||
https://github.com/user-attachments/assets/79af5fb8-707c-4263-b153-9ab2a01d3ac1
|
||||
|
||||
|
||||
|
||||
FastVideo currently offers: (with more to come)
|
||||
|
||||
- [NEW!] [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- FastHunyuan and FastMochi: consistency distilled video diffusion models for 8x inference speedup.
|
||||
- First open distillation recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
|
||||
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
|
||||
- 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.
|
||||
|
||||
|
||||
|
||||
## Change Log
|
||||
- ```2025/02/20```: FastVideo now supports STA on [StepVideo](https://github.com/stepfun-ai/Step-Video-T2V) with 3.4X speedup!
|
||||
- ```2025/02/18```: Release the inference code and kernel for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
- ```2025/01/13```: Support Lora finetuning for HunyuanVideo.
|
||||
- ```2024/12/25```: Enable single 4090 inference for `FastHunyuan`, please rerun the installation steps to update the environment.
|
||||
- ```2024/12/17```: `FastVideo` v1.0 is released.
|
||||
|
||||
|
||||
## 🔧 Installation
|
||||
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 [csrc/sliding_tile_attention/README.md](csrc/sliding_tile_attention/README.md) to install STA.
|
||||
|
||||
## 🚀 Inference
|
||||
### Inference StepVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
```
|
||||
python scripts/huggingface/download_hf.py --repo_id=stepfun-ai/stepvideo-t2v --local_dir=data/stepvideo-t2v --repo_type=model
|
||||
```
|
||||
Use the following scripts to run inference for StepVideo. When using STA for inference, the generated videos will have dimensions of 204×768×768 (currently, this is the only supported shape).
|
||||
```bash
|
||||
sh scripts/inference/inference_stepvideo_STA.sh # Inference stepvideo with STA
|
||||
sh scripts/inference/inference_stepvideo.sh # Inference original stepvideo
|
||||
```
|
||||
|
||||
### Inference HunyuanVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model
|
||||
```
|
||||
We provide two examples in the following script to run inference with STA + [TeaCache](https://github.com/ali-vilab/TeaCache) and STA only.
|
||||
```bash
|
||||
sh scripts/inference/inference_hunyuan_STA.sh
|
||||
```
|
||||
### Video Demos using STA + Teacache
|
||||
Visit our [demo website](https://fast-video.github.io/) to explore our complete collection of examples. We shorten a single video generation process from 945s to 317s on H100.
|
||||
|
||||
### Inference FastHunyuan on single RTX4090
|
||||
We now support NF4 and LLM-INT8 quantized inference using BitsAndBytes for FastHunyuan. With NF4 quantization, inference can be performed on a single RTX 4090 GPU, requiring just 20GB of VRAM.
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan-diffusers --local_dir=data/FastHunyuan-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan_hf_quantization.sh
|
||||
```
|
||||
For more information about the VRAM requirements for BitsAndBytes quantization, please refer to the table below (timing measured on an H100 GPU):
|
||||
|
||||
|
||||
| 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
|
||||
|
||||
```bash
|
||||
# Create and activate a new conda environment
|
||||
conda create -n fastvideo python=3.12
|
||||
conda activate fastvideo
|
||||
|
||||
# Install FastVideo
|
||||
pip install fastvideo
|
||||
# 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
|
||||
```
|
||||
|
||||
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) for more detailed installation instructions.
|
||||
|
||||
## Sparse Distillation
|
||||
For our sparse distillation techniques, please see our [distillation docs](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) and check out our [blog](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
|
||||
See below for recipes and datasets:
|
||||
|
||||
| Model | Sparse Distillation | Dataset |
|
||||
|:-------------------------------------------------------------------------------------------: |:---------------------------------------------------------------------------------------------------------------: |:--------------------------------------------------------------------------------------------------------: |
|
||||
| [FastWan2.1-T2V-1.3B](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P) | [FastVideo Synthetic Wan2.1 480P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k) |
|
||||
| [FastWan2.1-T2V-14B-Preview](https://huggingface.co/FastVideo/FastWan2.1-T2V-14B-Diffusers) | Coming soon! | [FastVideo Synthetic Wan2.1 720P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x768x1280_250k) |
|
||||
| [FastWan2.2-TI2V-5B](https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free) | [FastVideo Synthetic Wan2.2 720P](https://huggingface.co/datasets/FastVideo/Wan2.2-Syn-121x704x1280_32k) |
|
||||
|
||||
## Inference
|
||||
### Generating Your First Video
|
||||
Here's a minimal example to generate a video using the default settings. Make sure VSA kernels are [installed](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation.html). Create a file called `example.py` with the following code:
|
||||
|
||||
```python
|
||||
import os
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
def main():
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
|
||||
# Create a video generator with a pre-trained model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1, # Adjust based on your hardware
|
||||
)
|
||||
|
||||
# Define a prompt for your video
|
||||
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest."
|
||||
|
||||
# Generate the video
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
return_frames=True, # Also return frames from this call (defaults to False)
|
||||
output_path="my_videos/", # Controls where videos are saved
|
||||
save_video=True
|
||||
)
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
```
|
||||
|
||||
Run the script with:
|
||||
|
||||
## 🎯 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 example.py
|
||||
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):
|
||||
|
||||
For a more detailed guide, please see our [inference quick start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html).
|
||||
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
|
||||
|
||||
### Other docs:
|
||||
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.
|
||||
|
||||
- [Design Overview](https://hao-ai-lab.github.io/FastVideo/design/overview.html)
|
||||
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html)
|
||||
|
||||
## Distillation and Finetuning
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html)
|
||||
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
|
||||
Currently, both Mochi and Hunyuan models support Lora finetuning through diffusers. To generate personalized videos from your own dataset, you'll need to follow three main steps: dataset preparation, finetuning, and inference.
|
||||
|
||||
#### Dataset Preparation
|
||||
We provide scripts to better help you get started to train on your own characters!
|
||||
You can run this to organize your dataset to get the videos2caption.json before preprocess. Specify your video folder and corresponding caption folder (caption files should be .txt files and have the same name with its video):
|
||||
```
|
||||
python scripts/dataset_preparation/prepare_json_file.py --video_dir data/input_videos/ --prompt_dir data/captions/ --output_path data/output_folder/videos2caption.json --verbose
|
||||
```
|
||||
Also, we provide script to resize your videos:
|
||||
```
|
||||
python scripts/data_preprocess/resize_videos.py
|
||||
```
|
||||
#### Finetuning
|
||||
After basic dataset preparation and preprocess, you can start to finetune your model using Lora:
|
||||
```
|
||||
bash scripts/finetune/finetune_hunyuan_hf_lora.sh
|
||||
```
|
||||
#### Inference
|
||||
For inference with Lora checkpoint, you can run the following scripts with additional parameter `--lora_checkpoint_dir`:
|
||||
```
|
||||
bash scripts/inference/inference_hunyuan_hf.sh
|
||||
```
|
||||
**We also provide scripts for Mochi in the same directory.**
|
||||
|
||||
#### Finetune with Both Image and Video
|
||||
Our codebase support finetuning with both image and video.
|
||||
```bash
|
||||
bash scripts/finetune/finetune_hunyuan.sh
|
||||
bash scripts/finetune/finetune_mochi_lora_mix.sh
|
||||
```
|
||||
For Image-Video Mixture Fine-tuning, make sure to enable the `--group_frame` option in your script.
|
||||
|
||||
## 📑 Development Plan
|
||||
<!-- - More distillation methods -->
|
||||
<!-- - [ ] Add Distribution Matching Distillation -->
|
||||
More FastWan Models Coming Soon!
|
||||
- [ ] Add FastWan2.1-T2V-14B
|
||||
- [ ] Add FastWan2.2-T2V-14B
|
||||
- [ ] Add FastWan2.2-I2V-14B
|
||||
<!-- - Optimization features
|
||||
- Code updates -->
|
||||
<!-- - [ ] fp8 support -->
|
||||
<!-- - [ ] faster load model and save model support -->
|
||||
|
||||
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
|
||||
- More distillation methods
|
||||
- [ ] Add Distribution Matching Distillation
|
||||
- More models support
|
||||
- [ ] Add CogvideoX model
|
||||
- Code update
|
||||
- [ ] fp8 support
|
||||
- [ ] faster load model and save model support
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview.html)
|
||||
We welcome all contributions. Please run `bash format.sh --all` before submitting a pull request.
|
||||
|
||||
## 🔧 Testing
|
||||
Run `pytest` to verify the data preprocessing, checkpoint saving, and sequence parallel pipelines. We recommend adding corresponding test cases in the `test` folder to support your contribution.
|
||||
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects:
|
||||
- [Wan-Video](https://github.com/Wan-Video)
|
||||
- [ThunderKittens](https://github.com/HazyResearch/ThunderKittens)
|
||||
- [Triton](https://github.com/triton-lang/triton)
|
||||
- [DMD2](https://github.com/tianweiy/DMD2)
|
||||
- [diffusers](https://github.com/huggingface/diffusers)
|
||||
- [xDiT](https://github.com/xdit-project/xDiT)
|
||||
- [vLLM](https://github.com/vllm-project/vllm)
|
||||
- [SGLang](https://github.com/sgl-project/sglang)
|
||||
We learned and reused code from the following projects: [PCM](https://github.com/G-U-N/Phased-Consistency-Model), [diffusers](https://github.com/huggingface/diffusers), [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan), and [xDiT](https://github.com/xdit-project/xDiT).
|
||||
|
||||
We thank [MBZUAI](https://ifm.mbzuai.ac.ae/), [Anyscale](https://www.anyscale.com/), and [GMI Cloud](https://www.gmicloud.ai/) for their support throughout this project.
|
||||
We thank MBZUAI and Anyscale for their support throughout this project.
|
||||
|
||||
## Citation
|
||||
If you find FastVideo useful, please considering citing our work:
|
||||
## Citation
|
||||
If you use FastVideo for your research, please cite our paper:
|
||||
|
||||
```bibtex
|
||||
@software{fastvideo2024,
|
||||
title = {FastVideo: A Unified Framework for Accelerated Video Generation},
|
||||
author = {The FastVideo Team},
|
||||
url = {https://github.com/hao-ai-lab/FastVideo},
|
||||
month = apr,
|
||||
year = {2024},
|
||||
@misc{zhang2025fastvideogenerationsliding,
|
||||
title={Fast Video Generation with Sliding Tile Attention},
|
||||
author={Peiyuan Zhang and Yongqi Chen and Runlong Su and Hangliang Ding and Ion Stoica and Zhenghong Liu and Hao Zhang},
|
||||
year={2025},
|
||||
eprint={2502.04507},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV},
|
||||
url={https://arxiv.org/abs/2502.04507},
|
||||
}
|
||||
|
||||
@article{zhang2025vsa,
|
||||
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
|
||||
author={Zhang, Peiyuan and Huang, Haofeng and Chen, Yongqi and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
|
||||
journal={arXiv preprint arXiv:2505.13389},
|
||||
year={2025}
|
||||
}
|
||||
|
||||
@article{zhang2025fast,
|
||||
title={Fast video generation with sliding tile attention},
|
||||
author={Zhang, Peiyuan and Chen, Yongqi and Su, Runlong and Ding, Hangliang and Stoica, Ion and Liu, Zhengzhong and Zhang, Hao},
|
||||
journal={arXiv preprint arXiv:2502.04507},
|
||||
year={2025}
|
||||
@misc{ding2025efficientvditefficientvideodiffusion,
|
||||
title={Efficient-vDiT: Efficient Video Diffusion Transformers With Attention Tile},
|
||||
author={Hangliang Ding and Dacheng Li and Runlong Su and Peiyuan Zhang and Zhijie Deng and Ion Stoica and Hao Zhang},
|
||||
year={2025},
|
||||
eprint={2502.06155},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV},
|
||||
url={https://arxiv.org/abs/2502.06155},
|
||||
}
|
||||
```
|
||||
|
||||
@@ -1,15 +0,0 @@
|
||||
try:
|
||||
from .comfyui.video_generator.nodes import (NODE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS)
|
||||
WEB_DIRECTORY = "./web"
|
||||
__all__ = [
|
||||
'NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY'
|
||||
]
|
||||
except ImportError:
|
||||
# ComfyUI environment not available, skip comfyui imports
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
WEB_DIRECTORY = "./web"
|
||||
__all__ = [
|
||||
'NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY'
|
||||
]
|
||||
|
Before Width: | Height: | Size: 194 KiB |
|
After Width: | Height: | Size: 149 KiB |
@@ -1,6 +0,0 @@
|
||||
<svg width="160" height="93" viewBox="0 0 160 93" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<path d="M28.8511 91.66L57.6319 1.86368H64.5394L35.7585 91.66H28.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
<path d="M15.0376 91.66L43.8185 1.86368H46.1209L17.3401 91.66H15.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
<path d="M1.22217 91.66L30.003 1.86366H31.1543L2.3734 91.66H1.22217Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
|
||||
<path d="M71.4465 1.86483L42.666 91.6599H69.144L78.3538 58.2746H123.251L129.007 39.855H84.1099L89.866 22.5868H152.032L157.788 1.86483H71.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
</svg>
|
||||
|
Before Width: | Height: | Size: 691 B |
@@ -1,18 +0,0 @@
|
||||
<svg width="252" height="105" viewBox="0 0 252 105" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<path d="M89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457Z" fill="#356CFF"/>
|
||||
<path fill-rule="evenodd" clip-rule="evenodd" d="M96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM100.768 13.1217L87.7028 29.4852H103.143L100.768 13.1217Z" fill="#356CFF"/>
|
||||
<path d="M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057Z" fill="#356CFF"/>
|
||||
<path fill-rule="evenodd" clip-rule="evenodd" d="M108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM109.081 90.697L116.802 65.8487C116.802 65.8487 120.959 65.8487 132.242 65.8487C143.525 65.8487 137.586 78.5759 135.211 84.0304C133.307 88.4021 127.491 90.697 122.74 90.697C117.989 90.697 109.081 90.697 109.081 90.697Z" fill="#356CFF"/>
|
||||
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944H159.747C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273L125.188 48.273L124 37.97L147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852H131.836C120.142 29.4852 125.897 1.00043 141.337 1.00043L173.188 1.00056Z" fill="#356CFF"/>
|
||||
<path d="M179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056Z" fill="#356CFF"/>
|
||||
<path d="M161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457Z" fill="#356CFF"/>
|
||||
<path fill-rule="evenodd" clip-rule="evenodd" d="M230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM237.948 77.9692C239.984 70.6965 240.917 65.242 228.446 65.242C215.975 65.242 211.818 71.9087 210.037 77.9692C208.255 84.0298 208.255 91.3025 219.538 91.3025C230.821 91.3025 235.911 85.2419 237.948 77.9692Z" fill="#356CFF"/>
|
||||
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944M173.188 1.00056C173.188 1.00056 156.777 1.00043 141.337 1.00043M173.188 1.00056L141.337 1.00043M141.337 20.3944C146.088 20.3944 150.839 20.3944 159.747 20.3944M141.337 20.3944H159.747M159.747 20.3944C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273M148.463 48.273C139.556 48.273 125.188 48.273 125.188 48.273M148.463 48.273L125.188 48.273M125.188 48.273L124 37.97M124 37.97C124 37.97 141.931 37.97 147.87 37.97M124 37.97L147.87 37.97M147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852M151.433 29.4852C146.682 29.4852 138.962 29.4852 131.836 29.4852M151.433 29.4852H131.836M131.836 29.4852C120.142 29.4852 125.897 1.00043 141.337 1.00043M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057ZM96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM87.7028 29.4852L100.768 13.1217L103.143 29.4852H87.7028ZM89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457ZM108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM116.802 65.8487L109.081 90.697C109.081 90.697 117.989 90.697 122.74 90.697C127.491 90.697 133.307 88.4021 135.211 84.0304C137.586 78.5759 143.525 65.8487 132.242 65.8487C120.959 65.8487 116.802 65.8487 116.802 65.8487ZM179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056ZM161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457ZM230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM228.446 65.242C240.917 65.242 239.984 70.6965 237.948 77.9692C235.911 85.2419 230.821 91.3025 219.538 91.3025C208.255 91.3025 208.255 84.0298 210.037 77.9692C211.818 71.9087 215.975 65.242 228.446 65.242Z" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M15.2524 55.5451L21.191 100.999L24.7541 100.999L18.8156 55.5451L15.2524 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M8.12646 55.5451L14.065 100.999L15.2527 100.999L9.31417 55.5451L8.12646 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M1 55.5451L6.93853 100.999L7.53239 100.999L1.59385 55.5451L1 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
|
||||
<path d="M15.2524 48.2724L30.0988 1H33.6619L18.8156 48.2724H15.2524Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M8.12646 48.2724L22.9728 1H24.1605L9.31417 48.2724H8.12646Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M1 48.2724L15.8463 1H16.4402L1.59385 48.2724H1Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
|
||||
<path d="M85.3271 55.5457H67.5116L87 12.7363L44.3513 68.2729H58.6038L43.1636 101L85.3271 55.5457Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.18771" stroke-miterlimit="16"/>
|
||||
</svg>
|
||||
|
Before Width: | Height: | Size: 5.7 KiB |
|
Before Width: | Height: | Size: 303 KiB |
@@ -0,0 +1,24 @@
|
||||
# Configuration for Cog ⚙️
|
||||
# Reference: https://cog.run/yaml
|
||||
|
||||
build:
|
||||
gpu: true
|
||||
cuda: "12.1"
|
||||
python_version: "3.10"
|
||||
python_packages:
|
||||
- "torch==2.4.0"
|
||||
- "torchvision"
|
||||
- "ninja==1.11.1.3"
|
||||
- "transformers==4.46.1"
|
||||
- "git+https://github.com/huggingface/diffusers.git@bf64b32652a63a1865a0528a73a13652b201698b"
|
||||
- "accelerate==1.0.1"
|
||||
- "safetensors==0.4.5"
|
||||
- "peft==0.13.2"
|
||||
- "packaging==24.2"
|
||||
- "git+https://github.com/hao-ai-lab/FastVideo"
|
||||
|
||||
run:
|
||||
- FLASH_ATTENTION_SKIP_CUDA_BUILD=TRUE pip install flash-attn --no-build-isolation
|
||||
- curl -o /usr/local/bin/pget -L "https://github.com/replicate/pget/releases/latest/download/pget_$(uname -s)_$(uname -m)" && chmod +x /usr/local/bin/pget
|
||||
|
||||
predict: "predict.py:Predictor"
|
||||
@@ -1,777 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
# type: ignore
|
||||
# ruff: noqa
|
||||
# code borrowed from https://github.com/pytorch/pytorch/blob/main/torch/utils/collect_env.py
|
||||
# and vllm: https://github.com/vllm-project/vllm/blob/main/vllm/collect_env.py
|
||||
|
||||
import datetime
|
||||
import locale
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
import sys
|
||||
# Unlike the rest of the PyTorch this file must be python2 compliant.
|
||||
# This script outputs relevant system environment info
|
||||
# Run it with `python collect_env.py` or `python -m torch.utils.collect_env`
|
||||
from collections import namedtuple
|
||||
|
||||
from fastvideo.envs import environment_variables
|
||||
|
||||
try:
|
||||
import torch
|
||||
TORCH_AVAILABLE = True
|
||||
except (ImportError, NameError, AttributeError, OSError):
|
||||
TORCH_AVAILABLE = False
|
||||
|
||||
# System Environment Information
|
||||
SystemEnv = namedtuple(
|
||||
'SystemEnv',
|
||||
[
|
||||
'torch_version',
|
||||
'is_debug_build',
|
||||
'cuda_compiled_version',
|
||||
'gcc_version',
|
||||
'clang_version',
|
||||
'cmake_version',
|
||||
'os',
|
||||
'libc_version',
|
||||
'python_version',
|
||||
'python_platform',
|
||||
'is_cuda_available',
|
||||
'cuda_runtime_version',
|
||||
'cuda_module_loading',
|
||||
'nvidia_driver_version',
|
||||
'nvidia_gpu_models',
|
||||
'cudnn_version',
|
||||
'pip_version', # 'pip' or 'pip3'
|
||||
'pip_packages',
|
||||
'conda_packages',
|
||||
'hip_compiled_version',
|
||||
'hip_runtime_version',
|
||||
'miopen_runtime_version',
|
||||
'caching_allocator_config',
|
||||
'is_xnnpack_available',
|
||||
'cpu_info',
|
||||
'fastvideo_version',
|
||||
'fastvideo_build_flags',
|
||||
'gpu_topo',
|
||||
'env_vars',
|
||||
])
|
||||
|
||||
DEFAULT_CONDA_PATTERNS = {
|
||||
"torch",
|
||||
"numpy",
|
||||
"mypy"
|
||||
"cudatoolkit",
|
||||
"soumith",
|
||||
"mkl",
|
||||
"magma",
|
||||
"triton",
|
||||
"optree",
|
||||
"nccl",
|
||||
"transformers",
|
||||
"accelerate",
|
||||
"peft",
|
||||
"zmq",
|
||||
"nvidia",
|
||||
"pynvml",
|
||||
}
|
||||
|
||||
DEFAULT_PIP_PATTERNS = {
|
||||
"torch",
|
||||
"numpy",
|
||||
"mypy",
|
||||
"flake8",
|
||||
"triton",
|
||||
"optree",
|
||||
"onnx",
|
||||
"nccl",
|
||||
"transformers",
|
||||
"accelerate",
|
||||
"peft",
|
||||
"zmq",
|
||||
"nvidia",
|
||||
"pynvml",
|
||||
}
|
||||
|
||||
|
||||
def run(command):
|
||||
"""Return (return-code, stdout, stderr)."""
|
||||
shell = True if type(command) is str else False
|
||||
p = subprocess.Popen(command,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
shell=shell)
|
||||
raw_output, raw_err = p.communicate()
|
||||
rc = p.returncode
|
||||
if get_platform() == 'win32':
|
||||
enc = 'oem'
|
||||
else:
|
||||
enc = locale.getpreferredencoding()
|
||||
output = raw_output.decode(enc)
|
||||
if command == 'nvidia-smi topo -m':
|
||||
# don't remove the leading whitespace of `nvidia-smi topo -m`
|
||||
# because they are meaningful
|
||||
output = output.rstrip()
|
||||
else:
|
||||
output = output.strip()
|
||||
err = raw_err.decode(enc)
|
||||
return rc, output, err.strip()
|
||||
|
||||
|
||||
def run_and_read_all(run_lambda, command):
|
||||
"""Run command using run_lambda; reads and returns entire output if rc is 0."""
|
||||
rc, out, _ = run_lambda(command)
|
||||
if rc != 0:
|
||||
return None
|
||||
return out
|
||||
|
||||
|
||||
def run_and_parse_first_match(run_lambda, command, regex):
|
||||
"""Run command using run_lambda, returns the first regex match if it exists."""
|
||||
rc, out, _ = run_lambda(command)
|
||||
if rc != 0:
|
||||
return None
|
||||
match = re.search(regex, out)
|
||||
if match is None:
|
||||
return None
|
||||
return match.group(1)
|
||||
|
||||
|
||||
def run_and_return_first_line(run_lambda, command):
|
||||
"""Run command using run_lambda and returns first line if output is not empty."""
|
||||
rc, out, _ = run_lambda(command)
|
||||
if rc != 0:
|
||||
return None
|
||||
return out.split('\n')[0]
|
||||
|
||||
|
||||
def get_conda_packages(run_lambda, patterns=None):
|
||||
if patterns is None:
|
||||
patterns = DEFAULT_CONDA_PATTERNS
|
||||
conda = os.environ.get('CONDA_EXE', 'conda')
|
||||
out = run_and_read_all(run_lambda, "{} list".format(conda))
|
||||
if out is None:
|
||||
return out
|
||||
|
||||
return "\n".join(line for line in out.splitlines()
|
||||
if not line.startswith("#") and any(name in line
|
||||
for name in patterns))
|
||||
|
||||
|
||||
def get_gcc_version(run_lambda):
|
||||
return run_and_parse_first_match(run_lambda, 'gcc --version', r'gcc (.*)')
|
||||
|
||||
|
||||
def get_clang_version(run_lambda):
|
||||
return run_and_parse_first_match(run_lambda, 'clang --version',
|
||||
r'clang version (.*)')
|
||||
|
||||
|
||||
def get_cmake_version(run_lambda):
|
||||
return run_and_parse_first_match(run_lambda, 'cmake --version',
|
||||
r'cmake (.*)')
|
||||
|
||||
|
||||
def get_nvidia_driver_version(run_lambda):
|
||||
if get_platform() == 'darwin':
|
||||
cmd = 'kextstat | grep -i cuda'
|
||||
return run_and_parse_first_match(run_lambda, cmd,
|
||||
r'com[.]nvidia[.]CUDA [(](.*?)[)]')
|
||||
smi = get_nvidia_smi()
|
||||
return run_and_parse_first_match(run_lambda, smi, r'Driver Version: (.*?) ')
|
||||
|
||||
|
||||
def get_gpu_info(run_lambda):
|
||||
if get_platform() == 'darwin' or (TORCH_AVAILABLE and hasattr(
|
||||
torch.version, 'hip') and torch.version.hip is not None):
|
||||
if TORCH_AVAILABLE and torch.cuda.is_available():
|
||||
if torch.version.hip is not None:
|
||||
prop = torch.cuda.get_device_properties(0)
|
||||
if hasattr(prop, "gcnArchName"):
|
||||
gcnArch = " ({})".format(prop.gcnArchName)
|
||||
else:
|
||||
gcnArch = "NoGCNArchNameOnOldPyTorch"
|
||||
else:
|
||||
gcnArch = ""
|
||||
return torch.cuda.get_device_name(None) + gcnArch
|
||||
return None
|
||||
smi = get_nvidia_smi()
|
||||
uuid_regex = re.compile(r' \(UUID: .+?\)')
|
||||
rc, out, _ = run_lambda(smi + ' -L')
|
||||
if rc != 0:
|
||||
return None
|
||||
# Anonymize GPUs by removing their UUID
|
||||
return re.sub(uuid_regex, '', out)
|
||||
|
||||
|
||||
def get_running_cuda_version(run_lambda):
|
||||
return run_and_parse_first_match(run_lambda, 'nvcc --version',
|
||||
r'release .+ V(.*)')
|
||||
|
||||
|
||||
def get_cudnn_version(run_lambda):
|
||||
"""Return a list of libcudnn.so; it's hard to tell which one is being used."""
|
||||
if get_platform() == 'win32':
|
||||
system_root = os.environ.get('SYSTEMROOT', 'C:\\Windows')
|
||||
cuda_path = os.environ.get('CUDA_PATH', "%CUDA_PATH%")
|
||||
where_cmd = os.path.join(system_root, 'System32', 'where')
|
||||
cudnn_cmd = '{} /R "{}\\bin" cudnn*.dll'.format(where_cmd, cuda_path)
|
||||
elif get_platform() == 'darwin':
|
||||
# CUDA libraries and drivers can be found in /usr/local/cuda/. See
|
||||
# https://docs.nvidia.com/cuda/cuda-installation-guide-mac-os-x/index.html#install
|
||||
# https://docs.nvidia.com/deeplearning/sdk/cudnn-install/index.html#installmac
|
||||
# Use CUDNN_LIBRARY when cudnn library is installed elsewhere.
|
||||
cudnn_cmd = 'ls /usr/local/cuda/lib/libcudnn*'
|
||||
else:
|
||||
cudnn_cmd = 'ldconfig -p | grep libcudnn | rev | cut -d" " -f1 | rev'
|
||||
rc, out, _ = run_lambda(cudnn_cmd)
|
||||
# find will return 1 if there are permission errors or if not found
|
||||
if len(out) == 0 or (rc != 1 and rc != 0):
|
||||
l = os.environ.get('CUDNN_LIBRARY')
|
||||
if l is not None and os.path.isfile(l):
|
||||
return os.path.realpath(l)
|
||||
return None
|
||||
files_set = set()
|
||||
for fn in out.split('\n'):
|
||||
fn = os.path.realpath(fn) # eliminate symbolic links
|
||||
if os.path.isfile(fn):
|
||||
files_set.add(fn)
|
||||
if not files_set:
|
||||
return None
|
||||
# Alphabetize the result because the order is non-deterministic otherwise
|
||||
files = sorted(files_set)
|
||||
if len(files) == 1:
|
||||
return files[0]
|
||||
result = '\n'.join(files)
|
||||
return 'Probably one of the following:\n{}'.format(result)
|
||||
|
||||
|
||||
def get_nvidia_smi():
|
||||
# Note: nvidia-smi is currently available only on Windows and Linux
|
||||
smi = 'nvidia-smi'
|
||||
if get_platform() == 'win32':
|
||||
system_root = os.environ.get('SYSTEMROOT', 'C:\\Windows')
|
||||
program_files_root = os.environ.get('PROGRAMFILES', 'C:\\Program Files')
|
||||
legacy_path = os.path.join(program_files_root, 'NVIDIA Corporation',
|
||||
'NVSMI', smi)
|
||||
new_path = os.path.join(system_root, 'System32', smi)
|
||||
smis = [new_path, legacy_path]
|
||||
for candidate_smi in smis:
|
||||
if os.path.exists(candidate_smi):
|
||||
smi = '"{}"'.format(candidate_smi)
|
||||
break
|
||||
return smi
|
||||
|
||||
|
||||
def get_fastvideo_version():
|
||||
return ""
|
||||
from fastvideo import __version__, __version_tuple__
|
||||
|
||||
if __version__ == "dev":
|
||||
return "N/A (dev)"
|
||||
version_str = __version_tuple__[-1]
|
||||
if isinstance(version_str, str) and version_str.startswith('g'):
|
||||
# it's a dev build
|
||||
if '.' in version_str:
|
||||
# it's a dev build containing local changes
|
||||
git_sha = version_str.split('.')[0][1:]
|
||||
date = version_str.split('.')[-1][1:]
|
||||
return f"{__version__} (git sha: {git_sha}, date: {date})"
|
||||
else:
|
||||
# it's a dev build without local changes
|
||||
git_sha = version_str[1:] # type: ignore
|
||||
return f"{__version__} (git sha: {git_sha})"
|
||||
return __version__
|
||||
|
||||
|
||||
def summarize_fastvideo_build_flags():
|
||||
# This could be a static method if the flags are constant, or dynamic if you need to check environment variables, etc.
|
||||
return 'CUDA Archs: {}; ROCm: {}; Neuron: {}'.format(
|
||||
os.environ.get('TORCH_CUDA_ARCH_LIST', 'Not Set'),
|
||||
'Enabled' if os.environ.get('ROCM_HOME') else 'Disabled',
|
||||
'Enabled' if os.environ.get('NEURON_CORES') else 'Disabled',
|
||||
)
|
||||
|
||||
|
||||
def get_gpu_topo(run_lambda):
|
||||
output = None
|
||||
|
||||
if get_platform() == 'linux':
|
||||
output = run_and_read_all(run_lambda, 'nvidia-smi topo -m')
|
||||
if output is None:
|
||||
output = run_and_read_all(run_lambda, 'rocm-smi --showtopo')
|
||||
|
||||
return output
|
||||
|
||||
|
||||
# example outputs of CPU infos
|
||||
# * linux
|
||||
# Architecture: x86_64
|
||||
# CPU op-mode(s): 32-bit, 64-bit
|
||||
# Address sizes: 46 bits physical, 48 bits virtual
|
||||
# Byte Order: Little Endian
|
||||
# CPU(s): 128
|
||||
# On-line CPU(s) list: 0-127
|
||||
# Vendor ID: GenuineIntel
|
||||
# Model name: Intel(R) Xeon(R) Platinum 8375C CPU @ 2.90GHz
|
||||
# CPU family: 6
|
||||
# Model: 106
|
||||
# Thread(s) per core: 2
|
||||
# Core(s) per socket: 32
|
||||
# Socket(s): 2
|
||||
# Stepping: 6
|
||||
# BogoMIPS: 5799.78
|
||||
# Flags: fpu vme de pse tsc msr pae mce cx8 apic sep mtrr pge mca cmov pat pse36 clflush mmx fxsr
|
||||
# sse sse2 ss ht syscall nx pdpe1gb rdtscp lm constant_tsc arch_perfmon rep_good nopl
|
||||
# xtopology nonstop_tsc cpuid aperfmperf tsc_known_freq pni pclmulqdq monitor ssse3 fma cx16
|
||||
# pcid sse4_1 sse4_2 x2apic movbe popcnt tsc_deadline_timer aes xsave avx f16c rdrand
|
||||
# hypervisor lahf_lm abm 3dnowprefetch invpcid_single ssbd ibrs ibpb stibp ibrs_enhanced
|
||||
# fsgsbase tsc_adjust bmi1 avx2 smep bmi2 erms invpcid avx512f avx512dq rdseed adx smap
|
||||
# avx512ifma clflushopt clwb avx512cd sha_ni avx512bw avx512vl xsaveopt xsavec xgetbv1
|
||||
# xsaves wbnoinvd ida arat avx512vbmi pku ospke avx512_vbmi2 gfni vaes vpclmulqdq
|
||||
# avx512_vnni avx512_bitalg tme avx512_vpopcntdq rdpid md_clear flush_l1d arch_capabilities
|
||||
# Virtualization features:
|
||||
# Hypervisor vendor: KVM
|
||||
# Virtualization type: full
|
||||
# Caches (sum of all):
|
||||
# L1d: 3 MiB (64 instances)
|
||||
# L1i: 2 MiB (64 instances)
|
||||
# L2: 80 MiB (64 instances)
|
||||
# L3: 108 MiB (2 instances)
|
||||
# NUMA:
|
||||
# NUMA node(s): 2
|
||||
# NUMA node0 CPU(s): 0-31,64-95
|
||||
# NUMA node1 CPU(s): 32-63,96-127
|
||||
# Vulnerabilities:
|
||||
# Itlb multihit: Not affected
|
||||
# L1tf: Not affected
|
||||
# Mds: Not affected
|
||||
# Meltdown: Not affected
|
||||
# Mmio stale data: Vulnerable: Clear CPU buffers attempted, no microcode; SMT Host state unknown
|
||||
# Retbleed: Not affected
|
||||
# Spec store bypass: Mitigation; Speculative Store Bypass disabled via prctl and seccomp
|
||||
# Spectre v1: Mitigation; usercopy/swapgs barriers and __user pointer sanitization
|
||||
# Spectre v2: Mitigation; Enhanced IBRS, IBPB conditional, RSB filling, PBRSB-eIBRS SW sequence
|
||||
# Srbds: Not affected
|
||||
# Tsx async abort: Not affected
|
||||
# * win32
|
||||
# Architecture=9
|
||||
# CurrentClockSpeed=2900
|
||||
# DeviceID=CPU0
|
||||
# Family=179
|
||||
# L2CacheSize=40960
|
||||
# L2CacheSpeed=
|
||||
# Manufacturer=GenuineIntel
|
||||
# MaxClockSpeed=2900
|
||||
# Name=Intel(R) Xeon(R) Platinum 8375C CPU @ 2.90GHz
|
||||
# ProcessorType=3
|
||||
# Revision=27142
|
||||
#
|
||||
# Architecture=9
|
||||
# CurrentClockSpeed=2900
|
||||
# DeviceID=CPU1
|
||||
# Family=179
|
||||
# L2CacheSize=40960
|
||||
# L2CacheSpeed=
|
||||
# Manufacturer=GenuineIntel
|
||||
# MaxClockSpeed=2900
|
||||
# Name=Intel(R) Xeon(R) Platinum 8375C CPU @ 2.90GHz
|
||||
# ProcessorType=3
|
||||
# Revision=27142
|
||||
|
||||
|
||||
def get_cpu_info(run_lambda):
|
||||
rc, out, err = 0, '', ''
|
||||
if get_platform() == 'linux':
|
||||
rc, out, err = run_lambda('lscpu')
|
||||
elif get_platform() == 'win32':
|
||||
rc, out, err = run_lambda(
|
||||
'wmic cpu get Name,Manufacturer,Family,Architecture,ProcessorType,DeviceID, \
|
||||
CurrentClockSpeed,MaxClockSpeed,L2CacheSize,L2CacheSpeed,Revision /VALUE'
|
||||
)
|
||||
elif get_platform() == 'darwin':
|
||||
rc, out, err = run_lambda("sysctl -n machdep.cpu.brand_string")
|
||||
cpu_info = 'None'
|
||||
if rc == 0:
|
||||
cpu_info = out
|
||||
else:
|
||||
cpu_info = err
|
||||
return cpu_info
|
||||
|
||||
|
||||
def get_platform():
|
||||
if sys.platform.startswith('linux'):
|
||||
return 'linux'
|
||||
elif sys.platform.startswith('win32'):
|
||||
return 'win32'
|
||||
elif sys.platform.startswith('cygwin'):
|
||||
return 'cygwin'
|
||||
elif sys.platform.startswith('darwin'):
|
||||
return 'darwin'
|
||||
else:
|
||||
return sys.platform
|
||||
|
||||
|
||||
def get_mac_version(run_lambda):
|
||||
return run_and_parse_first_match(run_lambda, 'sw_vers -productVersion',
|
||||
r'(.*)')
|
||||
|
||||
|
||||
def get_windows_version(run_lambda):
|
||||
system_root = os.environ.get('SYSTEMROOT', 'C:\\Windows')
|
||||
wmic_cmd = os.path.join(system_root, 'System32', 'Wbem', 'wmic')
|
||||
findstr_cmd = os.path.join(system_root, 'System32', 'findstr')
|
||||
return run_and_read_all(
|
||||
run_lambda,
|
||||
'{} os get Caption | {} /v Caption'.format(wmic_cmd, findstr_cmd))
|
||||
|
||||
|
||||
def get_lsb_version(run_lambda):
|
||||
return run_and_parse_first_match(run_lambda, 'lsb_release -a',
|
||||
r'Description:\t(.*)')
|
||||
|
||||
|
||||
def check_release_file(run_lambda):
|
||||
return run_and_parse_first_match(run_lambda, 'cat /etc/*-release',
|
||||
r'PRETTY_NAME="(.*)"')
|
||||
|
||||
|
||||
def get_os(run_lambda):
|
||||
from platform import machine
|
||||
platform = get_platform()
|
||||
|
||||
if platform == 'win32' or platform == 'cygwin':
|
||||
return get_windows_version(run_lambda)
|
||||
|
||||
if platform == 'darwin':
|
||||
version = get_mac_version(run_lambda)
|
||||
if version is None:
|
||||
return None
|
||||
return 'macOS {} ({})'.format(version, machine())
|
||||
|
||||
if platform == 'linux':
|
||||
# Ubuntu/Debian based
|
||||
desc = get_lsb_version(run_lambda)
|
||||
if desc is not None:
|
||||
return '{} ({})'.format(desc, machine())
|
||||
|
||||
# Try reading /etc/*-release
|
||||
desc = check_release_file(run_lambda)
|
||||
if desc is not None:
|
||||
return '{} ({})'.format(desc, machine())
|
||||
|
||||
return '{} ({})'.format(platform, machine())
|
||||
|
||||
# Unknown platform
|
||||
return platform
|
||||
|
||||
|
||||
def get_python_platform():
|
||||
import platform
|
||||
return platform.platform()
|
||||
|
||||
|
||||
def get_libc_version():
|
||||
import platform
|
||||
if get_platform() != 'linux':
|
||||
return 'N/A'
|
||||
return '-'.join(platform.libc_ver())
|
||||
|
||||
|
||||
def get_pip_packages(run_lambda, patterns=None):
|
||||
"""Return `pip list` output. Note: will also find conda-installed pytorch and numpy packages."""
|
||||
if patterns is None:
|
||||
patterns = DEFAULT_PIP_PATTERNS
|
||||
|
||||
def run_with_pip():
|
||||
try:
|
||||
import importlib.util
|
||||
pip_spec = importlib.util.find_spec('pip')
|
||||
pip_available = pip_spec is not None
|
||||
except ImportError:
|
||||
pip_available = False
|
||||
|
||||
if pip_available:
|
||||
cmd = [sys.executable, '-mpip', 'list', '--format=freeze']
|
||||
elif os.environ.get("UV") is not None:
|
||||
print("uv is set")
|
||||
cmd = ["uv", "pip", "list", "--format=freeze"]
|
||||
else:
|
||||
raise RuntimeError(
|
||||
"Could not collect pip list output (pip or uv module not available)"
|
||||
)
|
||||
|
||||
out = run_and_read_all(run_lambda, cmd)
|
||||
return "\n".join(line for line in out.splitlines()
|
||||
if any(name in line for name in patterns))
|
||||
|
||||
pip_version = 'pip3' if sys.version[0] == '3' else 'pip'
|
||||
out = run_with_pip()
|
||||
return pip_version, out
|
||||
|
||||
|
||||
def get_cachingallocator_config():
|
||||
ca_config = os.environ.get('PYTORCH_CUDA_ALLOC_CONF', '')
|
||||
return ca_config
|
||||
|
||||
|
||||
def get_cuda_module_loading_config():
|
||||
if TORCH_AVAILABLE and torch.cuda.is_available():
|
||||
torch.cuda.init()
|
||||
config = os.environ.get('CUDA_MODULE_LOADING', '')
|
||||
return config
|
||||
else:
|
||||
return "N/A"
|
||||
|
||||
|
||||
def is_xnnpack_available():
|
||||
if TORCH_AVAILABLE:
|
||||
import torch.backends.xnnpack
|
||||
return str(torch.backends.xnnpack.enabled) # type: ignore[attr-defined]
|
||||
else:
|
||||
return "N/A"
|
||||
|
||||
|
||||
def get_env_vars():
|
||||
env_vars = ''
|
||||
secret_terms = ('secret', 'token', 'api', 'access', 'password')
|
||||
report_prefix = ("TORCH", "NCCL", "PYTORCH", "CUDA", "CUBLAS", "CUDNN",
|
||||
"OMP_", "MKL_", "NVIDIA")
|
||||
for k, v in os.environ.items():
|
||||
if any(term in k.lower() for term in secret_terms):
|
||||
continue
|
||||
if k in environment_variables:
|
||||
env_vars = env_vars + "{}={}".format(k, v) + "\n"
|
||||
if k.startswith(report_prefix):
|
||||
env_vars = env_vars + "{}={}".format(k, v) + "\n"
|
||||
|
||||
return env_vars
|
||||
|
||||
|
||||
def get_env_info():
|
||||
run_lambda = run
|
||||
pip_version, pip_list_output = get_pip_packages(run_lambda)
|
||||
|
||||
if TORCH_AVAILABLE:
|
||||
version_str = torch.__version__
|
||||
debug_mode_str = str(torch.version.debug)
|
||||
cuda_available_str = str(torch.cuda.is_available())
|
||||
cuda_version_str = torch.version.cuda
|
||||
if not hasattr(torch.version,
|
||||
'hip') or torch.version.hip is None: # cuda version
|
||||
hip_compiled_version = hip_runtime_version = miopen_runtime_version = 'N/A'
|
||||
else: # HIP version
|
||||
|
||||
def get_version_or_na(cfg, prefix):
|
||||
_lst = [s.rsplit(None, 1)[-1] for s in cfg if prefix in s]
|
||||
return _lst[0] if _lst else 'N/A'
|
||||
|
||||
cfg = torch._C._show_config().split('\n')
|
||||
hip_runtime_version = get_version_or_na(cfg, 'HIP Runtime')
|
||||
miopen_runtime_version = get_version_or_na(cfg, 'MIOpen')
|
||||
cuda_version_str = 'N/A'
|
||||
hip_compiled_version = torch.version.hip
|
||||
else:
|
||||
version_str = debug_mode_str = cuda_available_str = cuda_version_str = 'N/A'
|
||||
hip_compiled_version = hip_runtime_version = miopen_runtime_version = 'N/A'
|
||||
|
||||
sys_version = sys.version.replace("\n", " ")
|
||||
|
||||
conda_packages = get_conda_packages(run_lambda)
|
||||
|
||||
fastvideo_version = get_fastvideo_version()
|
||||
fastvideo_build_flags = summarize_fastvideo_build_flags()
|
||||
gpu_topo = get_gpu_topo(run_lambda)
|
||||
|
||||
return SystemEnv(
|
||||
torch_version=version_str,
|
||||
is_debug_build=debug_mode_str,
|
||||
python_version='{} ({}-bit runtime)'.format(
|
||||
sys_version,
|
||||
sys.maxsize.bit_length() + 1),
|
||||
python_platform=get_python_platform(),
|
||||
is_cuda_available=cuda_available_str,
|
||||
cuda_compiled_version=cuda_version_str,
|
||||
cuda_runtime_version=get_running_cuda_version(run_lambda),
|
||||
cuda_module_loading=get_cuda_module_loading_config(),
|
||||
nvidia_gpu_models=get_gpu_info(run_lambda),
|
||||
nvidia_driver_version=get_nvidia_driver_version(run_lambda),
|
||||
cudnn_version=get_cudnn_version(run_lambda),
|
||||
hip_compiled_version=hip_compiled_version,
|
||||
hip_runtime_version=hip_runtime_version,
|
||||
miopen_runtime_version=miopen_runtime_version,
|
||||
pip_version=pip_version,
|
||||
pip_packages=pip_list_output,
|
||||
conda_packages=conda_packages,
|
||||
os=get_os(run_lambda),
|
||||
libc_version=get_libc_version(),
|
||||
gcc_version=get_gcc_version(run_lambda),
|
||||
clang_version=get_clang_version(run_lambda),
|
||||
cmake_version=get_cmake_version(run_lambda),
|
||||
caching_allocator_config=get_cachingallocator_config(),
|
||||
is_xnnpack_available=is_xnnpack_available(),
|
||||
cpu_info=get_cpu_info(run_lambda),
|
||||
fastvideo_version=fastvideo_version,
|
||||
fastvideo_build_flags=fastvideo_build_flags,
|
||||
gpu_topo=gpu_topo,
|
||||
env_vars=get_env_vars(),
|
||||
)
|
||||
|
||||
|
||||
env_info_fmt = """
|
||||
PyTorch version: {torch_version}
|
||||
Is debug build: {is_debug_build}
|
||||
CUDA used to build PyTorch: {cuda_compiled_version}
|
||||
ROCM used to build PyTorch: {hip_compiled_version}
|
||||
|
||||
OS: {os}
|
||||
GCC version: {gcc_version}
|
||||
Clang version: {clang_version}
|
||||
CMake version: {cmake_version}
|
||||
Libc version: {libc_version}
|
||||
|
||||
Python version: {python_version}
|
||||
Python platform: {python_platform}
|
||||
Is CUDA available: {is_cuda_available}
|
||||
CUDA runtime version: {cuda_runtime_version}
|
||||
CUDA_MODULE_LOADING set to: {cuda_module_loading}
|
||||
GPU models and configuration: {nvidia_gpu_models}
|
||||
Nvidia driver version: {nvidia_driver_version}
|
||||
cuDNN version: {cudnn_version}
|
||||
HIP runtime version: {hip_runtime_version}
|
||||
MIOpen runtime version: {miopen_runtime_version}
|
||||
Is XNNPACK available: {is_xnnpack_available}
|
||||
|
||||
CPU:
|
||||
{cpu_info}
|
||||
|
||||
Versions of relevant libraries:
|
||||
{pip_packages}
|
||||
{conda_packages}
|
||||
""".strip()
|
||||
|
||||
# both the above code and the following code use `strip()` to
|
||||
# remove leading/trailing whitespaces, so we need to add a newline
|
||||
# in between to separate the two sections
|
||||
env_info_fmt += "\n"
|
||||
|
||||
env_info_fmt += """
|
||||
FastVideo Version: {fastvideo_version}
|
||||
FastVideo Build Flags:
|
||||
{fastvideo_build_flags}
|
||||
GPU Topology:
|
||||
{gpu_topo}
|
||||
|
||||
{env_vars}
|
||||
""".strip()
|
||||
|
||||
|
||||
def pretty_str(envinfo):
|
||||
|
||||
def replace_nones(dct, replacement='Could not collect'):
|
||||
for key in dct.keys():
|
||||
if dct[key] is not None:
|
||||
continue
|
||||
dct[key] = replacement
|
||||
return dct
|
||||
|
||||
def replace_bools(dct, true='Yes', false='No'):
|
||||
for key in dct.keys():
|
||||
if dct[key] is True:
|
||||
dct[key] = true
|
||||
elif dct[key] is False:
|
||||
dct[key] = false
|
||||
return dct
|
||||
|
||||
def prepend(text, tag='[prepend]'):
|
||||
lines = text.split('\n')
|
||||
updated_lines = [tag + line for line in lines]
|
||||
return '\n'.join(updated_lines)
|
||||
|
||||
def replace_if_empty(text, replacement='No relevant packages'):
|
||||
if text is not None and len(text) == 0:
|
||||
return replacement
|
||||
return text
|
||||
|
||||
def maybe_start_on_next_line(string):
|
||||
# If `string` is multiline, prepend a \n to it.
|
||||
if string is not None and len(string.split('\n')) > 1:
|
||||
return '\n{}\n'.format(string)
|
||||
return string
|
||||
|
||||
mutable_dict = envinfo._asdict()
|
||||
|
||||
# If nvidia_gpu_models is multiline, start on the next line
|
||||
mutable_dict['nvidia_gpu_models'] = \
|
||||
maybe_start_on_next_line(envinfo.nvidia_gpu_models)
|
||||
|
||||
# If the machine doesn't have CUDA, report some fields as 'No CUDA'
|
||||
dynamic_cuda_fields = [
|
||||
'cuda_runtime_version',
|
||||
'nvidia_gpu_models',
|
||||
'nvidia_driver_version',
|
||||
]
|
||||
all_cuda_fields = dynamic_cuda_fields + ['cudnn_version']
|
||||
all_dynamic_cuda_fields_missing = all(mutable_dict[field] is None
|
||||
for field in dynamic_cuda_fields)
|
||||
if TORCH_AVAILABLE and not torch.cuda.is_available(
|
||||
) and all_dynamic_cuda_fields_missing:
|
||||
for field in all_cuda_fields:
|
||||
mutable_dict[field] = 'No CUDA'
|
||||
if envinfo.cuda_compiled_version is None:
|
||||
mutable_dict['cuda_compiled_version'] = 'None'
|
||||
|
||||
# Replace True with Yes, False with No
|
||||
mutable_dict = replace_bools(mutable_dict)
|
||||
|
||||
# Replace all None objects with 'Could not collect'
|
||||
mutable_dict = replace_nones(mutable_dict)
|
||||
|
||||
# If either of these are '', replace with 'No relevant packages'
|
||||
mutable_dict['pip_packages'] = replace_if_empty(
|
||||
mutable_dict['pip_packages'])
|
||||
mutable_dict['conda_packages'] = replace_if_empty(
|
||||
mutable_dict['conda_packages'])
|
||||
|
||||
# Tag conda and pip packages with a prefix
|
||||
# If they were previously None, they'll show up as ie '[conda] Could not collect'
|
||||
if mutable_dict['pip_packages']:
|
||||
mutable_dict['pip_packages'] = prepend(
|
||||
mutable_dict['pip_packages'], '[{}] '.format(envinfo.pip_version))
|
||||
if mutable_dict['conda_packages']:
|
||||
mutable_dict['conda_packages'] = prepend(mutable_dict['conda_packages'],
|
||||
'[conda] ')
|
||||
mutable_dict['cpu_info'] = envinfo.cpu_info
|
||||
return env_info_fmt.format(**mutable_dict)
|
||||
|
||||
|
||||
def get_pretty_env_info():
|
||||
return pretty_str(get_env_info())
|
||||
|
||||
|
||||
def main():
|
||||
print("Collecting environment information...")
|
||||
output = get_pretty_env_info()
|
||||
print(output)
|
||||
|
||||
if TORCH_AVAILABLE and hasattr(torch, 'utils') and hasattr(
|
||||
torch.utils, '_crash_handler'):
|
||||
minidump_dir = torch.utils._crash_handler.DEFAULT_MINIDUMP_DIR
|
||||
if sys.platform == "linux" and os.path.exists(minidump_dir):
|
||||
dumps = [
|
||||
os.path.join(minidump_dir, dump)
|
||||
for dump in os.listdir(minidump_dir)
|
||||
]
|
||||
latest = max(dumps, key=os.path.getctime)
|
||||
ctime = os.path.getctime(latest)
|
||||
creation_time = datetime.datetime.fromtimestamp(ctime).strftime(
|
||||
'%Y-%m-%d %H:%M:%S')
|
||||
msg = "\n*** Detected a minidump at {} created on {}, ".format(latest, creation_time) + \
|
||||
"if this is related to your bug please include it when you file a report ***"
|
||||
print(msg, file=sys.stderr)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
main()
|
||||
@@ -1,138 +0,0 @@
|
||||
# ComfyUI-FastVideo
|
||||
|
||||
A custom node suite for ComfyUI that provides accelerated video generation using [FastVideo](https://github.com/hao-ai-labs/FastVideo). See the [blog post](https://hao-ai-lab.github.io/blogs/fastvideo/) about FastVideo V1 to learn more.
|
||||
|
||||
## Multi-GPU Parallel Inference
|
||||
|
||||
One of the key features ComfyUI-FastVideo brings to ComfyUI is its ability to distribute the generation workload across multiple GPUs, resulting in significantly faster inference times.
|
||||
|
||||

|
||||
|
||||
Example of Wan2.1-I2V-14B-480P-Diffusers model running on 4 GPUs.
|
||||
## Features
|
||||
|
||||
- Generate high-quality videos from text prompts and images
|
||||
- Configurable video parameters (prompt, resolution, frame count, FPS)
|
||||
- Support for multiple GPUs with tensor and sequence parallelism
|
||||
- Advanced configuration options for VAE, Text Encoder, and DIT components
|
||||
- Interruption/cancellation support for long-running generations
|
||||
|
||||
## Installation
|
||||
|
||||
### Requirements
|
||||
|
||||
- [ComfyUI](https://github.com/comfyanonymous/ComfyUI)
|
||||
- CUDA-capable GPU(s) with sufficient VRAM
|
||||
|
||||
### Install using ComfyUI Manager
|
||||
|
||||
Coming soon!
|
||||
|
||||
### Manual Installation
|
||||
|
||||
#### Copy the FastVideo `comfyui` directory into your ComfyUI custom_nodes directory:
|
||||
|
||||
```bash
|
||||
cp -r /path/to/FastVideo/comfyui /path/to/ComfyUI/custom_nodes/FastVideo
|
||||
```
|
||||
|
||||
#### Install dependencies:
|
||||
|
||||
Currently, the only dependency is `fastvideo`, which can be installed using pip.
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
```
|
||||
|
||||
#### Install missing custom nodes:
|
||||
|
||||
`ComfyUI-VideoHelperSuite`:
|
||||
|
||||
```bash
|
||||
cd /path/to/ComfyUI/custom_nodes
|
||||
git clone https://github.com/Kosinkadink/ComfyUI-VideoHelperSuite.git
|
||||
```
|
||||
|
||||
If you're seeing `ImportError: libGL.so.1: cannot open shared object file: No such file or directory`,
|
||||
you may need to install ffmpeg
|
||||
|
||||
```bash
|
||||
apt-get update && apt-get install ffmpeg
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
After installation, the following nodes will be available in the ComfyUI interface under the "fastvideo" category:
|
||||
|
||||
- **Video Generator**: The main node for generating videos from prompts
|
||||
- **Inference Args**: Configure video generation parameters
|
||||
- **VAE Config**
|
||||
- **Text Encoder Config**
|
||||
- **DIT Config**
|
||||
- **Load Image Path**: Load images for potential conditioning
|
||||
|
||||
You may have noticed many arguments on the nodes have 'auto' as the default value. This is because FastVideo will automatically detect the best values for these parameters based on the model and the hardware. However, you can also manually configure these parameters to get the best performance for your specific use case. We plan on releasing more optimized workflow files for different models and hardware configurations in the future.
|
||||
|
||||
You can see what some of the default configurations are by looking at the FastVideo repo:
|
||||
- [Wan2.1-I2V-14B-480P-Diffusers](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/configs/wan_14B_i2v_480p_pipeline.json)
|
||||
- [FastHunyuan-diffusers](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/configs/fasthunyuan_t2v.json)
|
||||
|
||||
### Node Configuration
|
||||
|
||||
#### Video Generator
|
||||
|
||||
- **prompt**: Text description of the video to generate
|
||||
- **output_path**: Directory where generated videos will be saved
|
||||
- **num_gpus**: Number of GPUs to use for generation
|
||||
- **model_path**: Path to the FastVideo model
|
||||
- **embedded_cfg_scale**: Classifier-free guidance scale
|
||||
- **sp_size**: Sequence parallelism size (usually should match num_gpus)
|
||||
- **tp_size**: Tensor parallelism size (usually should match num_gpus)
|
||||
- **precision**: Model precision (fp16 or bf16)
|
||||
|
||||
`model_path takes either a model id from huggingface or a local path to a model. Models by default will be downloaded to ~/.cache/huggingface/hub/ and cached for subsequent runs.`
|
||||
|
||||
#### Inference Args
|
||||
|
||||
- **height/width**: Resolution of the output video
|
||||
- **num_frames**: Number of frames to generate
|
||||
- **num_inference_steps**: Number of diffusion steps per frame
|
||||
- **guidance_scale**: Classifier-free guidance scale
|
||||
- **flow_shift**: Frame flow shift parameter
|
||||
- **seed**: Random seed for reproducible generation
|
||||
- **fps**: Frames per second of the output video
|
||||
- **image_path**: Optional path to input image for conditioning (for i2v models)
|
||||
|
||||
## Memory Management
|
||||
|
||||
Models will remain loaded in GPU memory between runs when you only change inference arguments (such as prompt, resolution, frame count, FPS, guidance scale, etc.) or the prompt text. This allows for faster subsequent generations since the model doesn't need to be reloaded.
|
||||
|
||||
However, if you need to change the following parameters, you will need to restart the ComfyUI server:
|
||||
- **Number of GPUs** (`num_gpus`)
|
||||
- **Model path** (`model_path`)
|
||||
- **Tensor parallelism size** (`tp_size`)
|
||||
- **Sequence parallelism size** (`sp_size`)
|
||||
|
||||
These parameters affect the model's distribution across GPUs and require a complete reinitialization of the model pipeline.
|
||||
|
||||
## Example workflows
|
||||
|
||||
### Text to Video
|
||||
|
||||
FastVideo-FastHunyuan-diffusers
|
||||
|
||||

|
||||
|
||||
- [FastHunyuan-diffusers.json](./examples/FastHunyuan-diffusers.json)
|
||||
|
||||
### Image to Video
|
||||
|
||||
Wan2.1-I2V-14B-480P-Diffusers
|
||||
|
||||

|
||||
|
||||
- [Wan2.1-I2V-14B-480P-Diffusers.json](./examples/Wan2.1-I2V-14B-480P-Diffusers.json)
|
||||
|
||||
## License
|
||||
|
||||
This project is licensed under Apache 2.0.
|
||||
@@ -1,5 +0,0 @@
|
||||
from .video_generator.nodes import (NODE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS)
|
||||
|
||||
WEB_DIRECTORY = "./web"
|
||||
__all__ = ['NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY']
|
||||
|
Before Width: | Height: | Size: 1.3 MiB |
@@ -1,6 +0,0 @@
|
||||
<svg width="160" height="93" viewBox="0 0 160 93" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<path d="M28.8511 91.66L57.6319 1.86368H64.5394L35.7585 91.66H28.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
<path d="M15.0376 91.66L43.8185 1.86368H46.1209L17.3401 91.66H15.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
<path d="M1.22217 91.66L30.003 1.86366H31.1543L2.3734 91.66H1.22217Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
|
||||
<path d="M71.4465 1.86483L42.666 91.6599H69.144L78.3538 58.2746H123.251L129.007 39.855H84.1099L89.866 22.5868H152.032L157.788 1.86483H71.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
</svg>
|
||||
|
Before Width: | Height: | Size: 691 B |
|
Before Width: | Height: | Size: 8.7 MiB |
|
Before Width: | Height: | Size: 769 KiB |
@@ -1,645 +0,0 @@
|
||||
{
|
||||
"id": "23a1f065-bbba-4a8f-b144-944e1318fcbf",
|
||||
"revision": 0,
|
||||
"last_node_id": 8,
|
||||
"last_link_id": 7,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 4,
|
||||
"type": "VAEConfig",
|
||||
"pos": [
|
||||
374.2159423828125,
|
||||
554.85888671875
|
||||
],
|
||||
"size": [
|
||||
334.080078125,
|
||||
322
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "vae_config",
|
||||
"type": "VAE_CONFIG",
|
||||
"links": [
|
||||
2
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VAEConfig"
|
||||
},
|
||||
"widgets_values": [
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"load_encoder": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
},
|
||||
"load_decoder": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
},
|
||||
"tile_sample_min_height": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 256
|
||||
},
|
||||
"tile_sample_min_width": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 256
|
||||
},
|
||||
"tile_sample_min_num_frames": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 16
|
||||
},
|
||||
"tile_sample_stride_height": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 192
|
||||
},
|
||||
"tile_sample_stride_width": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 192
|
||||
},
|
||||
"tile_sample_stride_num_frames": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 12
|
||||
},
|
||||
"blend_num_frames": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 0
|
||||
},
|
||||
"use_tiling": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
},
|
||||
"use_temporal_tiling": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
},
|
||||
"use_parallel_tiling": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "TextEncoderConfig",
|
||||
"pos": [
|
||||
416.4937744140625,
|
||||
953.6171875
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
106
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "text_encoder_config",
|
||||
"type": "TEXT_ENCODER_CONFIG",
|
||||
"links": [
|
||||
7
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "TextEncoderConfig"
|
||||
},
|
||||
"widgets_values": [
|
||||
-99999,
|
||||
-99999,
|
||||
-99999
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"prefix": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": ""
|
||||
},
|
||||
"quant_config": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": ""
|
||||
},
|
||||
"lora_config": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": ""
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 6,
|
||||
"type": "DITConfig",
|
||||
"pos": [
|
||||
415.1928405761719,
|
||||
1154.1573486328125
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
82
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "dit_config",
|
||||
"type": "DIT_CONFIG",
|
||||
"links": [
|
||||
6
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "DITConfig"
|
||||
},
|
||||
"widgets_values": [
|
||||
-99999,
|
||||
-99999
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"prefix": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": ""
|
||||
},
|
||||
"quant_config": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": ""
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 1,
|
||||
"type": "VideoGenerator",
|
||||
"pos": [
|
||||
818.804931640625,
|
||||
348.9299621582031
|
||||
],
|
||||
"size": [
|
||||
400,
|
||||
436
|
||||
],
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "inference_args",
|
||||
"shape": 7,
|
||||
"type": "INFERENCE_ARGS",
|
||||
"link": 3
|
||||
},
|
||||
{
|
||||
"name": "vae_config",
|
||||
"shape": 7,
|
||||
"type": "VAE_CONFIG",
|
||||
"link": 2
|
||||
},
|
||||
{
|
||||
"name": "text_encoder_config",
|
||||
"shape": 7,
|
||||
"type": "TEXT_ENCODER_CONFIG",
|
||||
"link": 7
|
||||
},
|
||||
{
|
||||
"name": "dit_config",
|
||||
"shape": 7,
|
||||
"type": "DIT_CONFIG",
|
||||
"link": 6
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "video_path",
|
||||
"type": "STRING",
|
||||
"links": [
|
||||
4
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VideoGenerator"
|
||||
},
|
||||
"widgets_values": [
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest. The playful yet serene atmosphere is complemented by soft natural light filtering through the petals. Mid-shot, warm and cheerful tones.",
|
||||
"/workspace/ComfyUI/outputs_video/",
|
||||
2,
|
||||
"FastVideo/FastHunyuan-diffusers",
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"embedded_cfg_scale": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 6
|
||||
},
|
||||
"sp_size": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 2
|
||||
},
|
||||
"tp_size": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 2
|
||||
},
|
||||
"vae_precision": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": "fp16"
|
||||
},
|
||||
"vae_tiling": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
},
|
||||
"vae_sp": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
},
|
||||
"text_encoder_precision": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": "fp16"
|
||||
},
|
||||
"precision": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": "fp16"
|
||||
},
|
||||
"dit_cpu_offload": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 3,
|
||||
"type": "VHS_LoadVideoPath",
|
||||
"pos": [
|
||||
1350.136962890625,
|
||||
331.20361328125
|
||||
],
|
||||
"size": [
|
||||
231.8896484375,
|
||||
286
|
||||
],
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"shape": 7,
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"shape": 7,
|
||||
"type": "VAE",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "video",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "video"
|
||||
},
|
||||
"link": 4
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
5
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "frame_count",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "AUDIO",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "video_info",
|
||||
"type": "VHS_VIDEOINFO",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VHS_LoadVideoPath"
|
||||
},
|
||||
"widgets_values": {
|
||||
"video": "",
|
||||
"force_rate": 0,
|
||||
"custom_width": 0,
|
||||
"custom_height": 0,
|
||||
"frame_load_cap": 0,
|
||||
"skip_first_frames": 0,
|
||||
"select_every_nth": 1,
|
||||
"format": "Wan",
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"filename": "",
|
||||
"type": "path",
|
||||
"format": "video/",
|
||||
"force_rate": 0,
|
||||
"custom_width": 0,
|
||||
"custom_height": 0,
|
||||
"frame_load_cap": 0,
|
||||
"skip_first_frames": 0,
|
||||
"select_every_nth": 1
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 2,
|
||||
"type": "InferenceArgs",
|
||||
"pos": [
|
||||
411.46307373046875,
|
||||
178.18182373046875
|
||||
],
|
||||
"size": [
|
||||
278.73828125,
|
||||
298
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "inference_args",
|
||||
"type": "INFERENCE_ARGS",
|
||||
"links": [
|
||||
3
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "InferenceArgs"
|
||||
},
|
||||
"widgets_values": [
|
||||
720,
|
||||
1280,
|
||||
45,
|
||||
6,
|
||||
-99999,
|
||||
-99999,
|
||||
1025,
|
||||
"fixed",
|
||||
24,
|
||||
-99999,
|
||||
-99999
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"height": {
|
||||
"isAuto": false,
|
||||
"value": 720,
|
||||
"cachedValue": 720
|
||||
},
|
||||
"width": {
|
||||
"isAuto": false,
|
||||
"value": 1280,
|
||||
"cachedValue": 1280
|
||||
},
|
||||
"num_frames": {
|
||||
"isAuto": false,
|
||||
"value": 45,
|
||||
"cachedValue": 45
|
||||
},
|
||||
"num_inference_steps": {
|
||||
"isAuto": false,
|
||||
"value": 6,
|
||||
"cachedValue": 6
|
||||
},
|
||||
"guidance_scale": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 1
|
||||
},
|
||||
"flow_shift": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 17
|
||||
},
|
||||
"seed": {
|
||||
"isAuto": false,
|
||||
"value": 1025,
|
||||
"cachedValue": 1024
|
||||
},
|
||||
"fps": {
|
||||
"isAuto": false,
|
||||
"value": 24,
|
||||
"cachedValue": 24
|
||||
},
|
||||
"image_path": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": "X://insert/path/here.mp4"
|
||||
},
|
||||
"enable_teacache": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 8,
|
||||
"type": "VHS_VideoCombine",
|
||||
"pos": [
|
||||
1668.3499755859375,
|
||||
328.22625732421875
|
||||
],
|
||||
"size": [
|
||||
507.507080078125,
|
||||
622.2227172851562
|
||||
],
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 5
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"shape": 7,
|
||||
"type": "AUDIO",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"shape": 7,
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"shape": 7,
|
||||
"type": "VAE",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "Filenames",
|
||||
"type": "VHS_FILENAMES",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VHS_VideoCombine"
|
||||
},
|
||||
"widgets_values": {
|
||||
"frame_rate": 24,
|
||||
"loop_count": 0,
|
||||
"filename_prefix": "",
|
||||
"format": "video/h264-mp4",
|
||||
"pix_fmt": "yuv420p",
|
||||
"crf": 19,
|
||||
"save_metadata": true,
|
||||
"trim_to_audio": false,
|
||||
"pingpong": false,
|
||||
"save_output": false,
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"filename": "._00003.mp4",
|
||||
"subfolder": "",
|
||||
"type": "temp",
|
||||
"format": "video/h264-mp4",
|
||||
"frame_rate": 24,
|
||||
"workflow": "._00003.png",
|
||||
"fullpath": "/workspace/ComfyUI/temp/._00003.mp4"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
2,
|
||||
4,
|
||||
0,
|
||||
1,
|
||||
1,
|
||||
"VAE_CONFIG"
|
||||
],
|
||||
[
|
||||
3,
|
||||
2,
|
||||
0,
|
||||
1,
|
||||
0,
|
||||
"INFERENCE_ARGS"
|
||||
],
|
||||
[
|
||||
4,
|
||||
1,
|
||||
0,
|
||||
3,
|
||||
2,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
5,
|
||||
3,
|
||||
0,
|
||||
8,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
6,
|
||||
6,
|
||||
0,
|
||||
1,
|
||||
3,
|
||||
"DIT_CONFIG"
|
||||
],
|
||||
[
|
||||
7,
|
||||
5,
|
||||
0,
|
||||
1,
|
||||
2,
|
||||
"TEXT_ENCODER_CONFIG"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.9090909090909091,
|
||||
"offset": [
|
||||
112.86678372727341,
|
||||
-71.45635903989245
|
||||
]
|
||||
},
|
||||
"frontendVersion": "1.20.4",
|
||||
"VHS_latentpreview": false,
|
||||
"VHS_latentpreviewrate": 0,
|
||||
"VHS_MetadataImage": true,
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
@@ -1,697 +0,0 @@
|
||||
{
|
||||
"id": "23a1f065-bbba-4a8f-b144-944e1318fcbf",
|
||||
"revision": 0,
|
||||
"last_node_id": 8,
|
||||
"last_link_id": 7,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 7,
|
||||
"type": "LoadImagePath",
|
||||
"pos": [
|
||||
33.15385437011719,
|
||||
191.2037353515625
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
334
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "image_path",
|
||||
"type": "STRING",
|
||||
"links": [
|
||||
1
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "MASK",
|
||||
"type": "MASK",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "LoadImagePath"
|
||||
},
|
||||
"widgets_values": [
|
||||
"woman.jpg",
|
||||
"image"
|
||||
]
|
||||
},
|
||||
{
|
||||
"id": 3,
|
||||
"type": "VHS_LoadVideoPath",
|
||||
"pos": [
|
||||
1350.136962890625,
|
||||
331.20361328125
|
||||
],
|
||||
"size": [
|
||||
231.8896484375,
|
||||
286
|
||||
],
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"shape": 7,
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"shape": 7,
|
||||
"type": "VAE",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "video",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "video"
|
||||
},
|
||||
"link": 4
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
5
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "frame_count",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "AUDIO",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "video_info",
|
||||
"type": "VHS_VIDEOINFO",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VHS_LoadVideoPath"
|
||||
},
|
||||
"widgets_values": {
|
||||
"video": "",
|
||||
"force_rate": 0,
|
||||
"custom_width": 0,
|
||||
"custom_height": 0,
|
||||
"frame_load_cap": 0,
|
||||
"skip_first_frames": 0,
|
||||
"select_every_nth": 1,
|
||||
"format": "Wan",
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"filename": "",
|
||||
"type": "path",
|
||||
"format": "video/",
|
||||
"force_rate": 0,
|
||||
"custom_width": 0,
|
||||
"custom_height": 0,
|
||||
"frame_load_cap": 0,
|
||||
"skip_first_frames": 0,
|
||||
"select_every_nth": 1
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 8,
|
||||
"type": "VHS_VideoCombine",
|
||||
"pos": [
|
||||
1668.3499755859375,
|
||||
328.22625732421875
|
||||
],
|
||||
"size": [
|
||||
214.7587890625,
|
||||
334
|
||||
],
|
||||
"flags": {},
|
||||
"order": 7,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 5
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"shape": 7,
|
||||
"type": "AUDIO",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"shape": 7,
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"shape": 7,
|
||||
"type": "VAE",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "Filenames",
|
||||
"type": "VHS_FILENAMES",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VHS_VideoCombine"
|
||||
},
|
||||
"widgets_values": {
|
||||
"frame_rate": 24,
|
||||
"loop_count": 0,
|
||||
"filename_prefix": "",
|
||||
"format": "video/h264-mp4",
|
||||
"pix_fmt": "yuv420p",
|
||||
"crf": 19,
|
||||
"save_metadata": true,
|
||||
"trim_to_audio": false,
|
||||
"pingpong": false,
|
||||
"save_output": false,
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {}
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 4,
|
||||
"type": "VAEConfig",
|
||||
"pos": [
|
||||
374.2159423828125,
|
||||
554.85888671875
|
||||
],
|
||||
"size": [
|
||||
334.080078125,
|
||||
322
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "vae_config",
|
||||
"type": "VAE_CONFIG",
|
||||
"links": [
|
||||
2
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VAEConfig"
|
||||
},
|
||||
"widgets_values": [
|
||||
true,
|
||||
true,
|
||||
256,
|
||||
256,
|
||||
16,
|
||||
192,
|
||||
192,
|
||||
12,
|
||||
0,
|
||||
true,
|
||||
true,
|
||||
true
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"load_encoder": {
|
||||
"isAuto": true,
|
||||
"value": true,
|
||||
"cachedValue": true
|
||||
},
|
||||
"load_decoder": {
|
||||
"isAuto": true,
|
||||
"value": true,
|
||||
"cachedValue": true
|
||||
},
|
||||
"tile_sample_min_height": {
|
||||
"isAuto": true,
|
||||
"value": 256,
|
||||
"cachedValue": 256
|
||||
},
|
||||
"tile_sample_min_width": {
|
||||
"isAuto": true,
|
||||
"value": 256,
|
||||
"cachedValue": 256
|
||||
},
|
||||
"tile_sample_min_num_frames": {
|
||||
"isAuto": true,
|
||||
"value": 16,
|
||||
"cachedValue": 16
|
||||
},
|
||||
"tile_sample_stride_height": {
|
||||
"isAuto": true,
|
||||
"value": 192,
|
||||
"cachedValue": 192
|
||||
},
|
||||
"tile_sample_stride_width": {
|
||||
"isAuto": true,
|
||||
"value": 192,
|
||||
"cachedValue": 192
|
||||
},
|
||||
"tile_sample_stride_num_frames": {
|
||||
"isAuto": true,
|
||||
"value": 12,
|
||||
"cachedValue": 12
|
||||
},
|
||||
"blend_num_frames": {
|
||||
"isAuto": true,
|
||||
"value": 0,
|
||||
"cachedValue": 0
|
||||
},
|
||||
"use_tiling": {
|
||||
"isAuto": true,
|
||||
"value": true,
|
||||
"cachedValue": true
|
||||
},
|
||||
"use_temporal_tiling": {
|
||||
"isAuto": true,
|
||||
"value": true,
|
||||
"cachedValue": true
|
||||
},
|
||||
"use_parallel_tiling": {
|
||||
"isAuto": true,
|
||||
"value": true,
|
||||
"cachedValue": true
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 2,
|
||||
"type": "InferenceArgs",
|
||||
"pos": [
|
||||
411.46307373046875,
|
||||
178.18182373046875
|
||||
],
|
||||
"size": [
|
||||
278.73828125,
|
||||
298
|
||||
],
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "image_path",
|
||||
"shape": 7,
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "image_path"
|
||||
},
|
||||
"link": 1
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "inference_args",
|
||||
"type": "INFERENCE_ARGS",
|
||||
"links": [
|
||||
3
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "InferenceArgs"
|
||||
},
|
||||
"widgets_values": [
|
||||
832,
|
||||
480,
|
||||
45,
|
||||
20,
|
||||
1,
|
||||
17,
|
||||
1024,
|
||||
"fixed",
|
||||
24,
|
||||
"X://insert/path/here.mp4",
|
||||
true
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"height": {
|
||||
"isAuto": false,
|
||||
"value": 832,
|
||||
"cachedValue": 720
|
||||
},
|
||||
"width": {
|
||||
"isAuto": false,
|
||||
"value": 480,
|
||||
"cachedValue": 1280
|
||||
},
|
||||
"num_frames": {
|
||||
"isAuto": false,
|
||||
"value": 45,
|
||||
"cachedValue": 45
|
||||
},
|
||||
"num_inference_steps": {
|
||||
"isAuto": false,
|
||||
"value": 20,
|
||||
"cachedValue": 6
|
||||
},
|
||||
"guidance_scale": {
|
||||
"isAuto": true,
|
||||
"value": 1,
|
||||
"cachedValue": 1
|
||||
},
|
||||
"flow_shift": {
|
||||
"isAuto": true,
|
||||
"value": 17,
|
||||
"cachedValue": 17
|
||||
},
|
||||
"seed": {
|
||||
"isAuto": false,
|
||||
"value": 1024,
|
||||
"cachedValue": 1024
|
||||
},
|
||||
"fps": {
|
||||
"isAuto": false,
|
||||
"value": 24,
|
||||
"cachedValue": 24
|
||||
},
|
||||
"image_path": {
|
||||
"isAuto": true,
|
||||
"value": "X://insert/path/here.mp4",
|
||||
"cachedValue": "X://insert/path/here.mp4"
|
||||
},
|
||||
"enable_teacache": {
|
||||
"isAuto": true,
|
||||
"value": true,
|
||||
"cachedValue": true
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 1,
|
||||
"type": "VideoGenerator",
|
||||
"pos": [
|
||||
818.804931640625,
|
||||
348.9299621582031
|
||||
],
|
||||
"size": [
|
||||
400,
|
||||
436
|
||||
],
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "inference_args",
|
||||
"shape": 7,
|
||||
"type": "INFERENCE_ARGS",
|
||||
"link": 3
|
||||
},
|
||||
{
|
||||
"name": "vae_config",
|
||||
"shape": 7,
|
||||
"type": "VAE_CONFIG",
|
||||
"link": 2
|
||||
},
|
||||
{
|
||||
"name": "text_encoder_config",
|
||||
"shape": 7,
|
||||
"type": "TEXT_ENCODER_CONFIG",
|
||||
"link": 7
|
||||
},
|
||||
{
|
||||
"name": "dit_config",
|
||||
"shape": 7,
|
||||
"type": "DIT_CONFIG",
|
||||
"link": 6
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "video_path",
|
||||
"type": "STRING",
|
||||
"links": [
|
||||
4
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VideoGenerator"
|
||||
},
|
||||
"widgets_values": [
|
||||
"A woman crying from laughter.",
|
||||
"/workspace/ComfyUI/outputs_video/",
|
||||
4,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
6,
|
||||
2,
|
||||
2,
|
||||
"fp16",
|
||||
true,
|
||||
true,
|
||||
"fp16",
|
||||
"fp16",
|
||||
true
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"embedded_cfg_scale": {
|
||||
"isAuto": true,
|
||||
"value": 6,
|
||||
"cachedValue": 6
|
||||
},
|
||||
"sp_size": {
|
||||
"isAuto": true,
|
||||
"value": 2,
|
||||
"cachedValue": 2
|
||||
},
|
||||
"tp_size": {
|
||||
"isAuto": true,
|
||||
"value": 2,
|
||||
"cachedValue": 2
|
||||
},
|
||||
"vae_precision": {
|
||||
"isAuto": true,
|
||||
"value": "fp16",
|
||||
"cachedValue": "fp16"
|
||||
},
|
||||
"vae_tiling": {
|
||||
"isAuto": true,
|
||||
"value": true,
|
||||
"cachedValue": true
|
||||
},
|
||||
"vae_sp": {
|
||||
"isAuto": true,
|
||||
"value": true,
|
||||
"cachedValue": true
|
||||
},
|
||||
"text_encoder_precision": {
|
||||
"isAuto": true,
|
||||
"value": "fp16",
|
||||
"cachedValue": "fp16"
|
||||
},
|
||||
"precision": {
|
||||
"isAuto": true,
|
||||
"value": "fp16",
|
||||
"cachedValue": "fp16"
|
||||
},
|
||||
"dit_cpu_offload": {
|
||||
"isAuto": true,
|
||||
"value": true,
|
||||
"cachedValue": true
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "TextEncoderConfig",
|
||||
"pos": [
|
||||
416.4937744140625,
|
||||
953.6171875
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
106
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "text_encoder_config",
|
||||
"type": "TEXT_ENCODER_CONFIG",
|
||||
"links": [
|
||||
7
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "TextEncoderConfig"
|
||||
},
|
||||
"widgets_values": [
|
||||
"",
|
||||
"",
|
||||
""
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"prefix": {
|
||||
"isAuto": true,
|
||||
"value": "",
|
||||
"cachedValue": ""
|
||||
},
|
||||
"quant_config": {
|
||||
"isAuto": true,
|
||||
"value": "",
|
||||
"cachedValue": ""
|
||||
},
|
||||
"lora_config": {
|
||||
"isAuto": true,
|
||||
"value": "",
|
||||
"cachedValue": ""
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 6,
|
||||
"type": "DITConfig",
|
||||
"pos": [
|
||||
415.1928405761719,
|
||||
1154.1573486328125
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
82
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "dit_config",
|
||||
"type": "DIT_CONFIG",
|
||||
"links": [
|
||||
6
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "DITConfig"
|
||||
},
|
||||
"widgets_values": [
|
||||
"",
|
||||
""
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"prefix": {
|
||||
"isAuto": true,
|
||||
"value": "",
|
||||
"cachedValue": ""
|
||||
},
|
||||
"quant_config": {
|
||||
"isAuto": true,
|
||||
"value": "",
|
||||
"cachedValue": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
1,
|
||||
7,
|
||||
0,
|
||||
2,
|
||||
0,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
2,
|
||||
4,
|
||||
0,
|
||||
1,
|
||||
1,
|
||||
"VAE_CONFIG"
|
||||
],
|
||||
[
|
||||
3,
|
||||
2,
|
||||
0,
|
||||
1,
|
||||
0,
|
||||
"INFERENCE_ARGS"
|
||||
],
|
||||
[
|
||||
4,
|
||||
1,
|
||||
0,
|
||||
3,
|
||||
2,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
5,
|
||||
3,
|
||||
0,
|
||||
8,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
6,
|
||||
6,
|
||||
0,
|
||||
1,
|
||||
3,
|
||||
"DIT_CONFIG"
|
||||
],
|
||||
[
|
||||
7,
|
||||
5,
|
||||
0,
|
||||
1,
|
||||
2,
|
||||
"TEXT_ENCODER_CONFIG"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.8264462809917354,
|
||||
"offset": [
|
||||
646.7950212991898,
|
||||
66.17259910028655
|
||||
]
|
||||
},
|
||||
"frontendVersion": "1.20.4",
|
||||
"VHS_latentpreview": false,
|
||||
"VHS_latentpreviewrate": 0,
|
||||
"VHS_MetadataImage": true,
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
@@ -1,31 +0,0 @@
|
||||
class DITConfig:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"optional": {
|
||||
"prefix": ("STRING", {
|
||||
"default": ""
|
||||
}),
|
||||
"quant_config": ("STRING", {
|
||||
"default": ""
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, **kwargs):
|
||||
return True
|
||||
|
||||
RETURN_TYPES = ("DIT_CONFIG", )
|
||||
RETURN_NAMES = ("dit_config", )
|
||||
FUNCTION = "set_args"
|
||||
CATEGORY = "fastvideo"
|
||||
|
||||
def set_args(self, prefix, quant_config):
|
||||
raw_args = {"prefix": prefix, "quant_config": quant_config}
|
||||
|
||||
# Filter out keys where value is -99999
|
||||
args = {k: v for k, v in raw_args.items() if str(int(v)) != str(-99999)}
|
||||
|
||||
return (args, )
|
||||
@@ -1,89 +0,0 @@
|
||||
class InferenceArgs:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"optional": {
|
||||
"height": ("INT", {
|
||||
"default": 720
|
||||
}),
|
||||
"width": ("INT", {
|
||||
"default": 1280
|
||||
}),
|
||||
"num_frames": ("INT", {
|
||||
"default": 45
|
||||
}),
|
||||
"num_inference_steps": ("INT", {
|
||||
"default": 6
|
||||
}),
|
||||
"guidance_scale": ("FLOAT", {
|
||||
"default": 1.0
|
||||
}),
|
||||
"flow_shift": ("INT", {
|
||||
"default": 17
|
||||
}),
|
||||
"seed": ("INT", {
|
||||
"default": 1024
|
||||
}),
|
||||
"fps": ("INT", {
|
||||
"default": 24
|
||||
}),
|
||||
"image_path": ("STRING", {
|
||||
"default": "X://insert/path/here.mp4"
|
||||
}),
|
||||
"enable_teacache": ([True, False], {
|
||||
"default": False
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, **kwargs):
|
||||
return True
|
||||
|
||||
RETURN_TYPES = ("INFERENCE_ARGS", )
|
||||
RETURN_NAMES = ("inference_args", )
|
||||
FUNCTION = "set_args"
|
||||
CATEGORY = "fastvideo"
|
||||
|
||||
def set_args(
|
||||
self,
|
||||
height,
|
||||
width,
|
||||
num_frames,
|
||||
num_inference_steps,
|
||||
guidance_scale,
|
||||
flow_shift,
|
||||
seed,
|
||||
fps,
|
||||
image_path,
|
||||
enable_teacache,
|
||||
):
|
||||
raw_args = {
|
||||
"height": height,
|
||||
"width": width,
|
||||
"num_frames": num_frames,
|
||||
"num_inference_steps": num_inference_steps,
|
||||
"guidance_scale": guidance_scale,
|
||||
"flow_shift": flow_shift,
|
||||
"seed": seed,
|
||||
"fps": fps,
|
||||
"image_path": image_path,
|
||||
"enable_teacache": enable_teacache,
|
||||
}
|
||||
|
||||
# Filter out keys where value is -99999, handling different types properly
|
||||
args = {}
|
||||
for k, v in raw_args.items():
|
||||
try:
|
||||
if isinstance(v, str):
|
||||
if v != "-99999":
|
||||
args[k] = v
|
||||
elif v != -99999:
|
||||
# If it's not a string, compare directly
|
||||
args[k] = v
|
||||
except (ValueError, TypeError):
|
||||
# Include any value that causes an error in comparison
|
||||
args[k] = v
|
||||
|
||||
return (args, )
|
||||
@@ -1,103 +0,0 @@
|
||||
import hashlib
|
||||
import os
|
||||
|
||||
import folder_paths
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image, ImageOps, ImageSequence
|
||||
|
||||
from .node_helpers import pillow
|
||||
|
||||
|
||||
class LoadImagePath:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
input_dir = folder_paths.get_input_directory()
|
||||
files = [
|
||||
f for f in os.listdir(input_dir)
|
||||
if os.path.isfile(os.path.join(input_dir, f))
|
||||
]
|
||||
files = folder_paths.filter_files_content_types(files, ["image"])
|
||||
return {
|
||||
"required": {
|
||||
"image": (sorted(files), {
|
||||
"image_upload": True
|
||||
})
|
||||
},
|
||||
}
|
||||
|
||||
CATEGORY = "fastvideo"
|
||||
|
||||
RETURN_TYPES = ("STRING", "IMAGE", "MASK")
|
||||
RETURN_NAMES = ("image_path", "IMAGE", "MASK")
|
||||
FUNCTION = "load_image"
|
||||
|
||||
def load_image(self, image):
|
||||
image_path = folder_paths.get_annotated_filepath(image)
|
||||
|
||||
img = pillow(Image.open, image_path)
|
||||
|
||||
output_images: list[torch.Tensor] = []
|
||||
output_masks: list[torch.Tensor] = []
|
||||
w, h = None, None
|
||||
|
||||
excluded_formats = ['MPO']
|
||||
|
||||
for i in ImageSequence.Iterator(img):
|
||||
processed_image = pillow(ImageOps.exif_transpose, i)
|
||||
if processed_image is None:
|
||||
continue
|
||||
|
||||
if processed_image.mode == 'I':
|
||||
processed_image = processed_image.point(lambda i: i * (1 / 255))
|
||||
image = processed_image.convert("RGB")
|
||||
|
||||
if len(output_images) == 0:
|
||||
w = image.size[0]
|
||||
h = image.size[1]
|
||||
|
||||
if image.size[0] != w or image.size[1] != h:
|
||||
continue
|
||||
|
||||
image = np.array(image).astype(np.float32) / 255.0
|
||||
image = torch.from_numpy(image)[
|
||||
None,
|
||||
]
|
||||
if 'A' in processed_image.getbands():
|
||||
mask = np.array(processed_image.getchannel('A')).astype(
|
||||
np.float32) / 255.0
|
||||
mask = 1. - torch.from_numpy(mask)
|
||||
elif processed_image.mode == 'P' and 'transparency' in processed_image.info:
|
||||
mask = np.array(
|
||||
processed_image.convert('RGBA').getchannel('A')).astype(
|
||||
np.float32) / 255.0
|
||||
mask = 1. - torch.from_numpy(mask)
|
||||
else:
|
||||
mask = torch.zeros((64, 64), dtype=torch.float32, device="cpu")
|
||||
output_images.append(image)
|
||||
output_masks.append(mask.unsqueeze(0))
|
||||
|
||||
if len(output_images) > 1 and img.format not in excluded_formats:
|
||||
output_image = torch.cat(output_images, dim=0)
|
||||
output_mask = torch.cat(output_masks, dim=0)
|
||||
else:
|
||||
output_image = output_images[0]
|
||||
output_mask = output_masks[0]
|
||||
|
||||
return (image_path, output_image, output_mask)
|
||||
|
||||
@classmethod
|
||||
def IS_CHANGED(s, image):
|
||||
image_path = folder_paths.get_annotated_filepath(image)
|
||||
m = hashlib.sha256()
|
||||
with open(image_path, 'rb') as f:
|
||||
m.update(f.read())
|
||||
return m.digest().hex()
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(s, image):
|
||||
if not folder_paths.exists_annotated_filepath(image):
|
||||
return "Invalid image file: {}".format(image)
|
||||
|
||||
return True
|
||||
@@ -1,68 +0,0 @@
|
||||
import hashlib
|
||||
from collections.abc import Callable
|
||||
from typing import Any, TypeVar
|
||||
|
||||
import torch
|
||||
from comfy.cli_args import args
|
||||
from PIL import ImageFile, UnidentifiedImageError
|
||||
|
||||
T = TypeVar('T')
|
||||
|
||||
|
||||
def conditioning_set_values(conditioning: list[Any],
|
||||
values: dict[str, Any] | None = None) -> list[Any]:
|
||||
if values is None:
|
||||
values = {}
|
||||
c = []
|
||||
for t in conditioning:
|
||||
n = [t[0], t[1].copy()]
|
||||
for k in values:
|
||||
n[1][k] = values[k]
|
||||
c.append(n)
|
||||
|
||||
return c
|
||||
|
||||
|
||||
def pillow(fn: Callable[[Any], T], arg: Any) -> T:
|
||||
prev_value = None
|
||||
try:
|
||||
x = fn(arg)
|
||||
except (OSError, UnidentifiedImageError, ValueError
|
||||
): #PIL issues #4472 and #2445, also fixes ComfyUI issue #3416
|
||||
prev_value = ImageFile.LOAD_TRUNCATED_IMAGES
|
||||
ImageFile.LOAD_TRUNCATED_IMAGES = True
|
||||
x = fn(arg)
|
||||
finally:
|
||||
if prev_value is not None:
|
||||
ImageFile.LOAD_TRUNCATED_IMAGES = prev_value
|
||||
return x
|
||||
|
||||
|
||||
def hasher() -> Callable[[], Any]:
|
||||
hashfuncs = {
|
||||
"md5": hashlib.md5,
|
||||
"sha1": hashlib.sha1,
|
||||
"sha256": hashlib.sha256,
|
||||
"sha512": hashlib.sha512
|
||||
}
|
||||
return hashfuncs[args.default_hashing_function]
|
||||
|
||||
|
||||
def string_to_torch_dtype(string: str) -> torch.dtype | None:
|
||||
if string == "fp32":
|
||||
return torch.float32
|
||||
if string == "fp16":
|
||||
return torch.float16
|
||||
if string == "bf16":
|
||||
return torch.bfloat16
|
||||
return None
|
||||
|
||||
|
||||
def image_alpha_fix(destination: torch.Tensor,
|
||||
source: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
if destination.shape[-1] < source.shape[-1]:
|
||||
source = source[..., :destination.shape[-1]]
|
||||
elif destination.shape[-1] > source.shape[-1]:
|
||||
destination = torch.nn.functional.pad(destination, (0, 1))
|
||||
destination[..., -1] = 1.0
|
||||
return destination, source
|
||||
@@ -1,24 +0,0 @@
|
||||
from .dit_config import DITConfig
|
||||
from .inference_args import InferenceArgs
|
||||
from .load_image import LoadImagePath
|
||||
from .text_encoder_config import TextEncoderConfig
|
||||
from .vae_config import VAEConfig
|
||||
from .video_generator import VideoGenerator
|
||||
|
||||
NODE_CLASS_MAPPINGS = {
|
||||
"VideoGenerator": VideoGenerator,
|
||||
"InferenceArgs": InferenceArgs,
|
||||
"VAEConfig": VAEConfig,
|
||||
"TextEncoderConfig": TextEncoderConfig,
|
||||
"DITConfig": DITConfig,
|
||||
"LoadImagePath": LoadImagePath
|
||||
}
|
||||
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"VideoGenerator": "Video Generator",
|
||||
"InferenceArgs": "Inference Args",
|
||||
"VAEConfig": "VAE Config",
|
||||
"TextEncoderConfig": "Text Encoder Config",
|
||||
"DITConfig": "DIT Config",
|
||||
"LoadImagePath": "Load Image Path"
|
||||
}
|
||||
@@ -1,38 +0,0 @@
|
||||
class TextEncoderConfig:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"optional": {
|
||||
"prefix": ("STRING", {
|
||||
"default": ""
|
||||
}),
|
||||
"quant_config": ("STRING", {
|
||||
"default": ""
|
||||
}),
|
||||
"lora_config": ("STRING", {
|
||||
"default": ""
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, **kwargs):
|
||||
return True
|
||||
|
||||
RETURN_TYPES = ("TEXT_ENCODER_CONFIG", )
|
||||
RETURN_NAMES = ("text_encoder_config", )
|
||||
FUNCTION = "set_args"
|
||||
CATEGORY = "fastvideo"
|
||||
|
||||
def set_args(self, prefix, quant_config, lora_config):
|
||||
raw_args = {
|
||||
"prefix": prefix,
|
||||
"quant_config": quant_config,
|
||||
"lora_config": lora_config
|
||||
}
|
||||
|
||||
# Filter out keys where value is -99999
|
||||
args = {k: v for k, v in raw_args.items() if str(int(v)) != str(-99999)}
|
||||
|
||||
return (args, )
|
||||
@@ -1,88 +0,0 @@
|
||||
class VAEConfig:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(cls):
|
||||
return {
|
||||
"optional": {
|
||||
"load_encoder": ([True, False], {
|
||||
"default": True
|
||||
}),
|
||||
"load_decoder": ([True, False], {
|
||||
"default": True
|
||||
}),
|
||||
"tile_sample_min_height": ("INT", {
|
||||
"default": 256
|
||||
}),
|
||||
"tile_sample_min_width": ("INT", {
|
||||
"default": 256
|
||||
}),
|
||||
"tile_sample_min_num_frames": ("INT", {
|
||||
"default": 16
|
||||
}),
|
||||
"tile_sample_stride_height": ("INT", {
|
||||
"default": 192
|
||||
}),
|
||||
"tile_sample_stride_width": ("INT", {
|
||||
"default": 192
|
||||
}),
|
||||
"tile_sample_stride_num_frames": ("INT", {
|
||||
"default": 12
|
||||
}),
|
||||
"blend_num_frames": ("INT", {
|
||||
"default": 0
|
||||
}),
|
||||
"use_tiling": ([True, False], {
|
||||
"default": True
|
||||
}),
|
||||
"use_temporal_tiling": ([True, False], {
|
||||
"default": True
|
||||
}),
|
||||
"use_parallel_tiling": ([True, False], {
|
||||
"default": True
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, **kwargs):
|
||||
return True
|
||||
|
||||
RETURN_TYPES = ("VAE_CONFIG", )
|
||||
RETURN_NAMES = ("vae_config", )
|
||||
FUNCTION = "set_args"
|
||||
CATEGORY = "fastvideo"
|
||||
|
||||
def set_args(
|
||||
self,
|
||||
load_encoder,
|
||||
load_decoder,
|
||||
tile_sample_min_height,
|
||||
tile_sample_min_width,
|
||||
tile_sample_min_num_frames,
|
||||
tile_sample_stride_height,
|
||||
tile_sample_stride_width,
|
||||
tile_sample_stride_num_frames,
|
||||
blend_num_frames,
|
||||
use_tiling,
|
||||
use_temporal_tiling,
|
||||
use_parallel_tiling,
|
||||
):
|
||||
raw_args = {
|
||||
"load_encoder": load_encoder,
|
||||
"load_decoder": load_decoder,
|
||||
"tile_sample_min_height": tile_sample_min_height,
|
||||
"tile_sample_min_width": tile_sample_min_width,
|
||||
"tile_sample_min_num_frames": tile_sample_min_num_frames,
|
||||
"tile_sample_stride_height": tile_sample_stride_height,
|
||||
"tile_sample_stride_width": tile_sample_stride_width,
|
||||
"tile_sample_stride_num_frames": tile_sample_stride_num_frames,
|
||||
"blend_num_frames": blend_num_frames,
|
||||
"use_tiling": use_tiling,
|
||||
"use_temporal_tiling": use_temporal_tiling,
|
||||
"use_parallel_tiling": use_parallel_tiling,
|
||||
}
|
||||
|
||||
# Filter out any value explicitly set to -99999
|
||||
args = {k: v for k, v in raw_args.items() if str(int(v)) != str(-99999)}
|
||||
|
||||
return (args, )
|
||||
@@ -1,315 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import glob
|
||||
import os
|
||||
import signal
|
||||
import sys
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from comfy.model_management import processing_interrupted
|
||||
|
||||
from fastvideo import PipelineConfig
|
||||
from fastvideo import VideoGenerator as FastVideoGenerator
|
||||
|
||||
sys.path.insert(
|
||||
0,
|
||||
os.path.dirname(
|
||||
os.path.dirname(
|
||||
os.path.dirname(os.path.dirname(os.path.abspath(__file__))))))
|
||||
|
||||
|
||||
# Custom exception for interruption
|
||||
class GenerationInterruptedException(Exception):
|
||||
pass
|
||||
|
||||
|
||||
# Custom exception for interruption that ComfyUI will recognize
|
||||
class GenerationCancelledException(Exception):
|
||||
|
||||
def __init__(self,
|
||||
message: str = "Generation was cancelled by user") -> None:
|
||||
self.message = message
|
||||
super().__init__(self.message)
|
||||
|
||||
|
||||
def update_config_from_args(config: Any, args_dict: dict[str, Any]) -> None:
|
||||
"""
|
||||
Update configuration object from arguments dictionary.
|
||||
|
||||
Args:
|
||||
config: The configuration object to update
|
||||
args_dict: Dictionary containing arguments
|
||||
"""
|
||||
for key, value in args_dict.items():
|
||||
if hasattr(config, key) and value is not None:
|
||||
if key == "text_encoder_precisions" and isinstance(value, list):
|
||||
setattr(config, key, tuple(value))
|
||||
else:
|
||||
setattr(config, key, value)
|
||||
|
||||
|
||||
class VideoGenerator:
|
||||
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {
|
||||
"required": {
|
||||
"prompt": ("STRING", {
|
||||
"multiline":
|
||||
True,
|
||||
"default":
|
||||
"A ripe orange tumbles gently from a tree and lands on the head of a lounging capybara, "
|
||||
"who blinks slowly in response. The moment is quietly humorous and oddly serene, framed by "
|
||||
"lush green foliage and dappled sunlight. Mid-shot, warm and whimsical tones."
|
||||
}),
|
||||
"output_path": ("STRING", {
|
||||
"default": "/workspace/ComfyUI/outputs_video/"
|
||||
}),
|
||||
"num_gpus": ("INT", {
|
||||
"default": 2,
|
||||
"min": 1,
|
||||
"max": 16
|
||||
}),
|
||||
"model_path": ("STRING", {
|
||||
"default": "FastVideo/FastHunyuan-diffusers"
|
||||
})
|
||||
},
|
||||
"optional": {
|
||||
"inference_args": ("INFERENCE_ARGS", ),
|
||||
"embedded_cfg_scale": ("FLOAT", {
|
||||
"default": 6.0
|
||||
}),
|
||||
"sp_size": ("INT", {
|
||||
"default": 2
|
||||
}),
|
||||
"tp_size": ("INT", {
|
||||
"default": 2
|
||||
}),
|
||||
"vae_config": ("VAE_CONFIG", ),
|
||||
"vae_precision": (["fp16", "bf16"], {
|
||||
"default": "fp16"
|
||||
}),
|
||||
"vae_tiling": ([True, False], {
|
||||
"default": True
|
||||
}),
|
||||
"vae_sp": ([True, False], {
|
||||
"default": False
|
||||
}),
|
||||
"text_encoder_config": ("TEXT_ENCODER_CONFIG", ),
|
||||
"text_encoder_precision": (["fp16", "bf16"], {
|
||||
"default": "fp16"
|
||||
}),
|
||||
"dit_config": ("DIT_CONFIG", ),
|
||||
"precision": (["fp16", "bf16"], {
|
||||
"default": "fp16"
|
||||
}),
|
||||
"dit_cpu_offload": ([True, False], {
|
||||
"default": False
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, **kwargs):
|
||||
return True
|
||||
|
||||
RETURN_TYPES = ("STRING", )
|
||||
RETURN_NAMES = ("video_path", )
|
||||
FUNCTION = "launch_inference"
|
||||
CATEGORY = "fastvideo"
|
||||
|
||||
generator: FastVideoGenerator | None = None
|
||||
_interrupt_thread: threading.Thread | None = None
|
||||
_generation_active: bool = False
|
||||
_generation_interrupted: bool = False
|
||||
_interrupt_event: threading.Event = threading.Event()
|
||||
_generation_thread: threading.Thread | None = None
|
||||
_generation_result: str | None = None
|
||||
_generation_exception: Exception | None = None
|
||||
|
||||
def _monitor_for_interruption(self):
|
||||
"""Background thread that monitors for interruption requests"""
|
||||
time.sleep(2) # Give the generation thread time to send execute_forward
|
||||
|
||||
while self._generation_active and not self._interrupt_event.is_set():
|
||||
if processing_interrupted():
|
||||
print("Video generation interrupted by user")
|
||||
self._generation_interrupted = True
|
||||
|
||||
# Try to send interrupt signal to worker processes
|
||||
if self.generator is not None and hasattr(
|
||||
self.generator, 'executor'):
|
||||
try:
|
||||
# The MultiprocExecutor has a workers attribute
|
||||
if hasattr(self.generator.executor, 'workers'):
|
||||
for worker in self.generator.executor.workers:
|
||||
if worker.is_alive():
|
||||
os.kill(worker.pid, signal.SIGINT)
|
||||
print("Interrupt signal sent to worker processes")
|
||||
except Exception as e:
|
||||
print(f"Error sending interrupt signal: {e}")
|
||||
|
||||
# Set the interrupt event to notify other threads
|
||||
self._interrupt_event.set()
|
||||
break
|
||||
time.sleep(0.5)
|
||||
|
||||
def _run_generation(self, prompt: str, output_path: str,
|
||||
inference_args: dict[str, Any]) -> None:
|
||||
"""Thread function to run the generation"""
|
||||
try:
|
||||
if self.generator is not None:
|
||||
self.generator.generate_video(prompt=prompt,
|
||||
output_path=output_path,
|
||||
**inference_args)
|
||||
self._generation_result = os.path.join(output_path,
|
||||
f"{prompt[:100]}.mp4")
|
||||
else:
|
||||
raise RuntimeError("Generator is not initialized")
|
||||
except Exception as e:
|
||||
self._generation_exception = e
|
||||
self._interrupt_event.set()
|
||||
|
||||
def load_output_video(self, output_dir):
|
||||
video_extensions = ["*.mp4", "*.avi", "*.mov", "*.mkv"]
|
||||
video_files = []
|
||||
|
||||
for ext in video_extensions:
|
||||
video_files.extend(glob.glob(os.path.join(output_dir, ext)))
|
||||
|
||||
if not video_files:
|
||||
print("No video files found in output directory: %s", output_dir)
|
||||
return ""
|
||||
|
||||
video_files.sort()
|
||||
return video_files[0]
|
||||
|
||||
def launch_inference(
|
||||
self,
|
||||
prompt,
|
||||
output_path,
|
||||
num_gpus,
|
||||
model_path,
|
||||
embedded_cfg_scale,
|
||||
sp_size,
|
||||
tp_size,
|
||||
vae_precision,
|
||||
vae_tiling,
|
||||
vae_sp,
|
||||
text_encoder_precision,
|
||||
precision,
|
||||
inference_args=None,
|
||||
vae_config=None,
|
||||
text_encoder_config=None,
|
||||
dit_config=None,
|
||||
dit_cpu_offload=None,
|
||||
):
|
||||
print('Running FastVideo inference')
|
||||
|
||||
# Reset interruption flag and event
|
||||
self._generation_interrupted = False
|
||||
self._interrupt_event.clear()
|
||||
self._generation_result = None
|
||||
self._generation_exception = None
|
||||
|
||||
# Load pipeline config from model path
|
||||
pipeline_config = PipelineConfig.from_pretrained(model_path)
|
||||
print('pipeline_config', pipeline_config)
|
||||
|
||||
# Update configs with provided config dictionaries
|
||||
if dit_config is not None:
|
||||
update_config_from_args(pipeline_config.dit_config, dit_config)
|
||||
|
||||
if vae_config is not None:
|
||||
update_config_from_args(pipeline_config.vae_config, vae_config)
|
||||
|
||||
if text_encoder_config is not None:
|
||||
update_config_from_args(pipeline_config.text_encoder_configs,
|
||||
text_encoder_config)
|
||||
|
||||
# Update top-level pipeline config with remaining arguments
|
||||
raw_pipeline_args = {}
|
||||
if embedded_cfg_scale is not None:
|
||||
raw_pipeline_args['embedded_cfg_scale'] = embedded_cfg_scale
|
||||
if precision is not None:
|
||||
raw_pipeline_args['precision'] = precision
|
||||
if vae_precision is not None:
|
||||
raw_pipeline_args['vae_precision'] = vae_precision
|
||||
if vae_tiling is not None:
|
||||
raw_pipeline_args['vae_tiling'] = vae_tiling
|
||||
if vae_sp is not None:
|
||||
raw_pipeline_args['vae_sp'] = vae_sp
|
||||
if text_encoder_precision is not None:
|
||||
raw_pipeline_args['text_encoder_precision'] = text_encoder_precision
|
||||
|
||||
# Filter out any value explicitly set to -99999 (auto values)
|
||||
pipeline_args = {
|
||||
k: v
|
||||
for k, v in raw_pipeline_args.items() if str(int(v)) != str(-99999)
|
||||
}
|
||||
|
||||
update_config_from_args(pipeline_config, pipeline_args)
|
||||
|
||||
raw_generation_args = {}
|
||||
if num_gpus is not None:
|
||||
raw_generation_args['num_gpus'] = num_gpus
|
||||
if tp_size is not None:
|
||||
raw_generation_args['tp_size'] = tp_size
|
||||
if sp_size is not None:
|
||||
raw_generation_args['sp_size'] = sp_size
|
||||
if dit_cpu_offload is not None:
|
||||
raw_generation_args['dit_cpu_offload'] = dit_cpu_offload
|
||||
|
||||
generation_args = {
|
||||
k: v
|
||||
for k, v in raw_generation_args.items()
|
||||
if str(int(v)) != str(-99999)
|
||||
}
|
||||
|
||||
if self.generator is None:
|
||||
print('generation_args', generation_args)
|
||||
print('pipeline_config', pipeline_config)
|
||||
self.generator = FastVideoGenerator.from_pretrained(
|
||||
model_path=model_path,
|
||||
**generation_args,
|
||||
pipeline_config=pipeline_config)
|
||||
|
||||
print('inference_args', inference_args)
|
||||
|
||||
# Start a thread to run the generation
|
||||
self._generation_thread = threading.Thread(target=self._run_generation,
|
||||
args=(prompt, output_path,
|
||||
inference_args),
|
||||
daemon=True)
|
||||
self._generation_thread.start()
|
||||
|
||||
# Start a background thread to monitor for interruptions
|
||||
self._generation_active = True
|
||||
self._interrupt_thread = threading.Thread(
|
||||
target=self._monitor_for_interruption, daemon=True)
|
||||
self._interrupt_thread.start()
|
||||
|
||||
# Wait for either completion or interruption
|
||||
while self._generation_thread.is_alive(
|
||||
) and not self._interrupt_event.is_set():
|
||||
self._generation_thread.join(timeout=0.5)
|
||||
|
||||
self._generation_active = False
|
||||
if self._interrupt_thread:
|
||||
self._interrupt_thread.join(timeout=1.0)
|
||||
self._interrupt_thread = None
|
||||
|
||||
if self._generation_interrupted:
|
||||
print("Video generation was cancelled by user")
|
||||
raise GenerationCancelledException()
|
||||
elif self._generation_exception:
|
||||
# Re-raise the exception from the generation thread
|
||||
raise self._generation_exception
|
||||
elif self._generation_result:
|
||||
return (self._generation_result, )
|
||||
else:
|
||||
# This shouldn't happen, but just in case
|
||||
print("Generation completed but no result was produced")
|
||||
raise Exception("Generation failed to produce a result")
|
||||
@@ -1,593 +0,0 @@
|
||||
import { app } from '../../../scripts/app.js'
|
||||
|
||||
function chainCallback(object, property, callback) {
|
||||
if (object == undefined) {
|
||||
console.error("Tried to add callback to non-existent object");
|
||||
return;
|
||||
}
|
||||
if (property in object && object[property]) {
|
||||
const callback_orig = object[property];
|
||||
object[property] = function () {
|
||||
const r = callback_orig.apply(this, arguments);
|
||||
return callback.apply(this, arguments) ?? r;
|
||||
};
|
||||
} else {
|
||||
object[property] = callback;
|
||||
}
|
||||
}
|
||||
|
||||
function drawAutoAnnotated(ctx, node, widget_width, y, H) {
|
||||
const litegraph_base = LiteGraph;
|
||||
const show_text = app.canvas.ds.scale >= 0.5;
|
||||
const margin = 15;
|
||||
|
||||
const autoTextWidth = 30;
|
||||
const autoTextRightMargin = 5;
|
||||
|
||||
ctx.textAlign = 'left';
|
||||
ctx.strokeStyle = litegraph_base.WIDGET_OUTLINE_COLOR;
|
||||
ctx.fillStyle = litegraph_base.WIDGET_BGCOLOR;
|
||||
|
||||
ctx.beginPath();
|
||||
if (show_text && ctx.roundRect) {
|
||||
ctx.roundRect(margin, y, widget_width - margin * 2, H, [H * 0.5]);
|
||||
} else {
|
||||
ctx.rect(margin, y, widget_width - margin * 2, H);
|
||||
}
|
||||
ctx.fill();
|
||||
|
||||
if (show_text) {
|
||||
if (!this.disabled) ctx.stroke();
|
||||
const isAuto = this.isAuto === true;
|
||||
|
||||
ctx.save();
|
||||
if (isAuto) {
|
||||
ctx.fillStyle = litegraph_base.WIDGET_TEXT_COLOR;
|
||||
ctx.strokeStyle = litegraph_base.WIDGET_TEXT_COLOR;
|
||||
} else {
|
||||
ctx.fillStyle = litegraph_base.WIDGET_SECONDARY_TEXT_COLOR;
|
||||
ctx.strokeStyle = litegraph_base.WIDGET_SECONDARY_TEXT_COLOR;
|
||||
}
|
||||
|
||||
// Position for the cog
|
||||
const cogX = widget_width - autoTextRightMargin - autoTextWidth - 6;
|
||||
const cogY = y + H * 0.5;
|
||||
const cogRadius = 6;
|
||||
const toothLength = 2;
|
||||
const numTeeth = 8;
|
||||
const holeRadius = 2; // Radius of the center hole
|
||||
|
||||
// Draw the cog
|
||||
ctx.beginPath();
|
||||
ctx.arc(cogX, cogY, cogRadius - toothLength, 0, Math.PI * 2);
|
||||
ctx.fill();
|
||||
|
||||
// Draw the center hole (by clearing it)
|
||||
ctx.beginPath();
|
||||
ctx.arc(cogX, cogY, holeRadius, 0, Math.PI * 2);
|
||||
ctx.fillStyle = litegraph_base.WIDGET_BGCOLOR;
|
||||
ctx.fill();
|
||||
|
||||
// Reset fill style for the teeth
|
||||
if (isAuto) {
|
||||
ctx.fillStyle = litegraph_base.WIDGET_TEXT_COLOR;
|
||||
} else {
|
||||
ctx.fillStyle = litegraph_base.WIDGET_SECONDARY_TEXT_COLOR;
|
||||
}
|
||||
|
||||
// Draw teeth
|
||||
ctx.beginPath();
|
||||
for (let i = 0; i < numTeeth; i++) {
|
||||
const angle = (i / numTeeth) * Math.PI * 2;
|
||||
const innerX = cogX + (cogRadius - toothLength) * Math.cos(angle);
|
||||
const innerY = cogY + (cogRadius - toothLength) * Math.sin(angle);
|
||||
const outerX = cogX + cogRadius * Math.cos(angle);
|
||||
const outerY = cogY + cogRadius * Math.sin(angle);
|
||||
|
||||
ctx.moveTo(innerX, innerY);
|
||||
ctx.lineTo(outerX, outerY);
|
||||
}
|
||||
ctx.lineWidth = 2;
|
||||
ctx.stroke();
|
||||
|
||||
ctx.restore();
|
||||
|
||||
// Draw label
|
||||
ctx.fillStyle = litegraph_base.WIDGET_SECONDARY_TEXT_COLOR;
|
||||
const label = this.label || this.name;
|
||||
if (label != null) {
|
||||
ctx.fillText(label, margin * 2 + 5, y + H * 0.7);
|
||||
}
|
||||
|
||||
// Draw value
|
||||
ctx.textAlign = 'right';
|
||||
const text = isAuto ? "auto" : this.displayValue();
|
||||
ctx.fillStyle = isAuto ? litegraph_base.WIDGET_SECONDARY_TEXT_COLOR : litegraph_base.WIDGET_TEXT_COLOR;
|
||||
ctx.fillText(text, widget_width - autoTextRightMargin - autoTextWidth - 15, y + H * 0.7);
|
||||
|
||||
// Draw increment/decrement buttons if not in AUTO mode and not a string widget
|
||||
if (!isAuto && !this.disabled && this.config[0] !== "FVAUTOSTRING") {
|
||||
// Draw decrement button (left triangle)
|
||||
ctx.fillStyle = litegraph_base.WIDGET_TEXT_COLOR;
|
||||
ctx.beginPath();
|
||||
ctx.moveTo(margin + 16, y + 5);
|
||||
ctx.lineTo(margin + 6, y + H * 0.5);
|
||||
ctx.lineTo(margin + 16, y + H - 5);
|
||||
ctx.fill();
|
||||
|
||||
// Draw increment button (right triangle)
|
||||
ctx.beginPath();
|
||||
ctx.moveTo(widget_width - margin - 16, y + 5);
|
||||
ctx.lineTo(widget_width - margin - 6, y + H * 0.5);
|
||||
ctx.lineTo(widget_width - margin - 16, y + H - 5);
|
||||
ctx.fill();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
function mouseAutoAnnotated(event, [x, y], node) {
|
||||
const widget_width = node.size[0];
|
||||
const margin = 15;
|
||||
const H = 20; // Widget height
|
||||
|
||||
const autoTextWidth = 30;
|
||||
const autoTextRightMargin = 5;
|
||||
|
||||
const cogRadius = 6;
|
||||
|
||||
if (this.isAuto) {
|
||||
if (event.type === "pointerup" || event.type === "mouseup") {
|
||||
const cogX = widget_width - autoTextRightMargin - autoTextWidth - 6;
|
||||
const cogLeftEdge = cogX - cogRadius;
|
||||
const cogRightEdge = cogX + cogRadius;
|
||||
|
||||
if (x > cogLeftEdge && x < cogRightEdge) {
|
||||
this.isAuto = false;
|
||||
this.value = this.cachedValue !== undefined ? this.cachedValue : (this.options.default || 0);
|
||||
|
||||
if (this.callback) {
|
||||
this.callback(this.value);
|
||||
}
|
||||
node.graph.setDirtyCanvas(true, false);
|
||||
}
|
||||
}
|
||||
|
||||
// Block ALL events in auto mode except cog clicks
|
||||
event.preventDefault?.();
|
||||
event.stopPropagation?.();
|
||||
event.stopImmediatePropagation?.();
|
||||
return true; // Always return true to indicate event was handled
|
||||
}
|
||||
|
||||
// Determine if clicking on increment/decrement buttons
|
||||
const delta = this.config[0] === "FVAUTOSTRING" ? 0 :
|
||||
(x < 40 ? -1 : x > widget_width - 48 ? 1 : 0);
|
||||
|
||||
if (event.type === "pointerdown" || event.type === "mousedown") {
|
||||
// ComfyUI appears to intercept pointerdown events, so this code path is never reached
|
||||
console.log("pointerdown received (unexpected)");
|
||||
return false;
|
||||
} else if (event.type === "pointerup" || event.type === "mouseup") {
|
||||
// Stop event propagation to prevent double handling
|
||||
event.preventDefault?.();
|
||||
event.stopPropagation?.();
|
||||
event.stopImmediatePropagation?.();
|
||||
|
||||
const cogX = widget_width - autoTextRightMargin - autoTextWidth - 6;
|
||||
const cogLeftEdge = widget_width - autoTextRightMargin - autoTextWidth - 6 - cogRadius;
|
||||
const cogRightEdge = widget_width - autoTextRightMargin - autoTextWidth - 6 + cogRadius;
|
||||
|
||||
if (x > cogLeftEdge && x < cogRightEdge) {
|
||||
this.isAuto = !this.isAuto;
|
||||
|
||||
if (this.isAuto) {
|
||||
this.cachedValue = this.value;
|
||||
this.value = -99999;
|
||||
} else {
|
||||
this.value = this.cachedValue !== undefined ? this.cachedValue : (this.options.default || 0);
|
||||
}
|
||||
|
||||
if (this.callback) {
|
||||
this.callback(this.value);
|
||||
}
|
||||
|
||||
node.graph.setDirtyCanvas(true, false);
|
||||
return true;
|
||||
}
|
||||
|
||||
// If in auto mode and NOT clicking the cog, block all other interactions
|
||||
if (this.isAuto) {
|
||||
return true;
|
||||
}
|
||||
|
||||
// Handle increment/decrement buttons if not in auto mode
|
||||
if (delta !== 0 && !this.isAuto) {
|
||||
if (this.config[0] === "FVAUTOCOMBO") {
|
||||
const options = this.options.values || [];
|
||||
if (options.length === 0) return true;
|
||||
|
||||
let currentIndex = -1;
|
||||
for (let i = 0; i < options.length; i++) {
|
||||
const optValue = typeof options[i] === 'object' ? options[i].value : options[i];
|
||||
if (optValue == this.value || String(optValue) === String(this.value)) {
|
||||
currentIndex = i;
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
if (currentIndex === -1) {
|
||||
currentIndex = 0;
|
||||
}
|
||||
|
||||
let newIndex = currentIndex + delta;
|
||||
if (newIndex < 0) {
|
||||
newIndex = options.length - 1;
|
||||
} else if (newIndex >= options.length) {
|
||||
newIndex = 0;
|
||||
}
|
||||
|
||||
const newOption = options[newIndex];
|
||||
this.value = typeof newOption === 'object' ? newOption.value : newOption;
|
||||
|
||||
if (this.callback) {
|
||||
this.callback(this.value);
|
||||
}
|
||||
|
||||
node.graph.setDirtyCanvas(true, false);
|
||||
return true;
|
||||
} else {
|
||||
let v = parseFloat(this.value);
|
||||
const increment = delta * 0.1 * (this.options.step || 1);
|
||||
|
||||
v += increment;
|
||||
|
||||
// Apply min/max constraints
|
||||
if (this.options.min != null) {
|
||||
v = Math.max(this.options.min, v);
|
||||
}
|
||||
if (this.options.max != null) {
|
||||
v = Math.min(this.options.max, v);
|
||||
}
|
||||
|
||||
// Round to precision or to integer
|
||||
if (this.config[0] === "FVAUTOINT") {
|
||||
v = Math.round(v);
|
||||
} else if (this.options.precision !== undefined) {
|
||||
const precision = Math.pow(10, this.options.precision);
|
||||
v = Math.round(v * precision) / precision;
|
||||
}
|
||||
|
||||
this.value = v;
|
||||
|
||||
if (this.callback) {
|
||||
this.callback(this.value);
|
||||
}
|
||||
|
||||
node.graph.setDirtyCanvas(true, false);
|
||||
return true;
|
||||
}
|
||||
}
|
||||
|
||||
if (delta === 0 && !this.isAuto) {
|
||||
if (this.config[0] === "FVAUTOCOMBO") {
|
||||
const options = this.options.values || [];
|
||||
|
||||
// Create menu items
|
||||
const menuItems = options.map(opt => {
|
||||
const value = typeof opt === 'object' ? opt.value : opt;
|
||||
const label = typeof opt === 'object' ? opt.label : opt.toString();
|
||||
|
||||
return {
|
||||
content: label,
|
||||
callback: () => {
|
||||
this.value = value;
|
||||
|
||||
if (this.callback) {
|
||||
this.callback(this.value);
|
||||
}
|
||||
|
||||
node.graph.setDirtyCanvas(true, false);
|
||||
}
|
||||
};
|
||||
});
|
||||
|
||||
new LiteGraph.ContextMenu(menuItems, {
|
||||
event: event,
|
||||
title: null,
|
||||
callback: null,
|
||||
extra: node
|
||||
});
|
||||
|
||||
return true;
|
||||
} else if (this.config[0] === "FVAUTOSTRING") {
|
||||
const d_callback = (v) => {
|
||||
this.value = v;
|
||||
|
||||
if (this.callback) {
|
||||
this.callback(this.value);
|
||||
}
|
||||
|
||||
node.graph.setDirtyCanvas(true, false);
|
||||
};
|
||||
|
||||
const dialog = app.canvas.prompt(
|
||||
'Value',
|
||||
this.value,
|
||||
d_callback,
|
||||
event
|
||||
);
|
||||
|
||||
return true;
|
||||
} else {
|
||||
// For numeric widgets, show input dialog
|
||||
const d_callback = (v) => {
|
||||
this.value = this.parseValue?.(v) ?? Number(v);
|
||||
|
||||
// Apply min/max constraints
|
||||
if (this.options.min != null) {
|
||||
this.value = Math.max(this.options.min, this.value);
|
||||
}
|
||||
if (this.options.max != null) {
|
||||
this.value = Math.min(this.options.max, this.value);
|
||||
}
|
||||
|
||||
// Round to precision or to integer
|
||||
if (this.config[0] === "FVAUTOINT") {
|
||||
this.value = Math.round(this.value);
|
||||
} else if (this.options.precision !== undefined) {
|
||||
const precision = Math.pow(10, this.options.precision);
|
||||
this.value = Math.round(this.value * precision) / precision;
|
||||
}
|
||||
|
||||
if (this.callback) {
|
||||
this.callback(this.value);
|
||||
}
|
||||
|
||||
node.graph.setDirtyCanvas(true, false);
|
||||
};
|
||||
|
||||
const dialog = app.canvas.prompt(
|
||||
'Value',
|
||||
this.value,
|
||||
d_callback,
|
||||
event
|
||||
);
|
||||
|
||||
return true;
|
||||
}
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
return false;
|
||||
}
|
||||
|
||||
function makeAutoAnnotated(widget, inputData) {
|
||||
const original = {
|
||||
callback: widget.callback,
|
||||
type: widget.type,
|
||||
value: widget.value
|
||||
};
|
||||
|
||||
// Add AUTO properties to the widget
|
||||
Object.assign(widget, {
|
||||
type: "BOOLEAN",
|
||||
draw: drawAutoAnnotated,
|
||||
mouse: mouseAutoAnnotated,
|
||||
onMouse: null, // Explicitly disable original onMouse handler
|
||||
|
||||
isAuto: true,
|
||||
cachedValue: widget.value,
|
||||
config: inputData,
|
||||
options: Object.assign({}, inputData[1], widget.options),
|
||||
original: original, // Store original properties for reference
|
||||
|
||||
// Disable other potential mouse handlers with no-op functions
|
||||
onClick: function () {
|
||||
return false;
|
||||
},
|
||||
onPointerUp: function () {
|
||||
return false;
|
||||
},
|
||||
onPointerDown: function () {
|
||||
return false;
|
||||
},
|
||||
onMouseUp: function () {
|
||||
return false;
|
||||
},
|
||||
onMouseDown: function () {
|
||||
return false;
|
||||
},
|
||||
|
||||
computeSize(width) {
|
||||
return [width, 20];
|
||||
},
|
||||
displayValue: function () {
|
||||
if (this.config[0] === "FVAUTOINT") {
|
||||
return Math.round(this.value).toString();
|
||||
}
|
||||
if (this.config[0] === "FVAUTOCOMBO") {
|
||||
return this.value;
|
||||
}
|
||||
if (this.config[0] === "FVAUTOSTRING") {
|
||||
return this.value;
|
||||
}
|
||||
// For FLOAT values, check if it's actually an integer
|
||||
if (Number.isInteger(this.value)) {
|
||||
return this.value.toString();
|
||||
}
|
||||
return this.value.toFixed(this.options.precision || 2);
|
||||
},
|
||||
parseValue: function (v) {
|
||||
if (this.config[0] === "FVAUTOSTRING") {
|
||||
return v;
|
||||
}
|
||||
if (typeof v === "string") {
|
||||
return parseFloat(v);
|
||||
}
|
||||
return v;
|
||||
},
|
||||
serializeValue: function () {
|
||||
// Return special value for AUTO mode
|
||||
return this.isAuto ? -99999 : this.value;
|
||||
},
|
||||
deserializeValue: function (data) {
|
||||
if (data === -99999) {
|
||||
this.isAuto = true;
|
||||
this.value = -99999;
|
||||
} else {
|
||||
this.isAuto = false;
|
||||
this.value = data;
|
||||
this.cachedValue = data;
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
// Override callback to handle AUTO mode
|
||||
widget.callback = function (v) {
|
||||
if (this.isAuto) {
|
||||
return; // Don't call the original callback in AUTO mode
|
||||
}
|
||||
const result = original.callback?.call(this, v);
|
||||
return result;
|
||||
};
|
||||
|
||||
// Override any potential click handlers
|
||||
const originalOnClick = widget.onClick;
|
||||
if (originalOnClick) {
|
||||
widget.onClick = function (...args) {
|
||||
if (this.isAuto) {
|
||||
return false;
|
||||
}
|
||||
return originalOnClick.call(this, ...args);
|
||||
};
|
||||
}
|
||||
|
||||
return widget;
|
||||
}
|
||||
|
||||
app.registerExtension({
|
||||
name: "FastVideo.AutoWidgets",
|
||||
|
||||
async beforeRegisterNodeDef(nodeType, nodeData, app) {
|
||||
if (nodeData?.name == "VideoGenerator" || nodeData?.name === "InferenceArgs" || nodeData?.name === "VAEConfig" ||
|
||||
nodeData?.name === "TextEncoderConfig" || nodeData?.name === "DITConfig") {
|
||||
// Add serialization support
|
||||
chainCallback(nodeType.prototype, "onSerialize", function (info) {
|
||||
if (!this.widgets) {
|
||||
return;
|
||||
}
|
||||
// Ensure widgets_values exists
|
||||
if (!info.widgets_values) {
|
||||
info.widgets_values = {};
|
||||
}
|
||||
|
||||
// Store AUTO widget states in a separate property
|
||||
if (!info.auto_widget_states) {
|
||||
info.auto_widget_states = {};
|
||||
}
|
||||
|
||||
// Handle AUTO widgets specially
|
||||
for (const w of this.widgets) {
|
||||
if (w.type === "BOOLEAN" && w.isAuto !== undefined) {
|
||||
// Store the serialized value (for Python node)
|
||||
info.widgets_values[w.name] = w.serializeValue();
|
||||
|
||||
// Store the full state (for UI restoration)
|
||||
info.auto_widget_states[w.name] = {
|
||||
isAuto: w.isAuto,
|
||||
value: w.value,
|
||||
cachedValue: w.cachedValue
|
||||
};
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
// Add deserialization support
|
||||
chainCallback(nodeType.prototype, "onConfigure", function (info) {
|
||||
if (!this.widgets) {
|
||||
return;
|
||||
}
|
||||
|
||||
// First, restore from widgets_values (for backward compatibility)
|
||||
if (info.widgets_values && Array.isArray(info.widgets_values)) {
|
||||
for (let i = 0; i < this.widgets.length && i < info.widgets_values.length; i++) {
|
||||
const w = this.widgets[i];
|
||||
const value = info.widgets_values[i];
|
||||
|
||||
if (w.type === "BOOLEAN" && w.isAuto !== undefined) {
|
||||
w.deserializeValue(value);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Then, restore full state if available
|
||||
if (info.auto_widget_states) {
|
||||
for (const w of this.widgets) {
|
||||
if (w.type === "BOOLEAN" && w.isAuto !== undefined && w.name in info.auto_widget_states) {
|
||||
const state = info.auto_widget_states[w.name];
|
||||
|
||||
w.isAuto = state.isAuto;
|
||||
w.cachedValue = state.cachedValue;
|
||||
w.value = state.isAuto ? -99999 : state.value;
|
||||
|
||||
w.callback?.(w.value);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Force a redraw
|
||||
this.graph?.setDirtyCanvas(true, true);
|
||||
});
|
||||
|
||||
|
||||
// Override addInput to handle AUTO widgets
|
||||
chainCallback(nodeType.prototype, "onNodeCreated", function () {
|
||||
// Convert any existing widgets to AUTO widgets if needed
|
||||
let new_widgets = [];
|
||||
const intWidgetNames = ["sp_size", "tp_size", "height", "width", "num_frames", "num_inference_steps", "flow_shift", "seed", "fps", "scale_factor",
|
||||
"tile_sample_min_height", "tile_sample_min_width", "tile_sample_min_num_frames", "tile_sample_stride_height", "tile_sample_stride_width",
|
||||
"tile_sample_stride_num_frames", "blend_num_frames"
|
||||
]
|
||||
const floatWidgetNames = ["embedded_cfg_scale", "guidance_scale"]
|
||||
const comboWidgetNames = ["vae_tiling", "vae_precision", "vae_sp", "text_encoder_precision", "precision",
|
||||
"load_encoder", "load_decoder", "use_tiling", "use_temporal_tiling", "use_parallel_tiling", "dit_cpu_offload", "enable_teacache"
|
||||
]
|
||||
const stringWidgetNames = ["prefix", "quant_config", "lora_config", "image_path"]
|
||||
|
||||
if (this.widgets) {
|
||||
for (let w of this.widgets) {
|
||||
if (intWidgetNames.includes(w.name)) {
|
||||
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOINT", { "default": 0 }]));
|
||||
} else if (floatWidgetNames.includes(w.name)) {
|
||||
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOFLOAT", { "default": 0 }]));
|
||||
} else if (comboWidgetNames.includes(w.name)) {
|
||||
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOCOMBO", { "default": 0 }]));
|
||||
} else if (stringWidgetNames.includes(w.name)) {
|
||||
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOSTRING", { "default": "" }]));
|
||||
} else {
|
||||
new_widgets.push(w);
|
||||
}
|
||||
}
|
||||
this.widgets = new_widgets;
|
||||
|
||||
const autoWidgets = this.widgets.filter(w => w.type === "BOOLEAN" && w.isAuto !== undefined);
|
||||
}
|
||||
|
||||
this.graph?.setDirtyCanvas(true, true);
|
||||
});
|
||||
}
|
||||
},
|
||||
|
||||
async init() {
|
||||
// Force a redraw of all nodes when the extension initializes
|
||||
if (app.graph) {
|
||||
setTimeout(() => {
|
||||
app.graph.setDirtyCanvas(true, true);
|
||||
}, 1000);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
console.log("FastVideo.core.js loaded");
|
||||
@@ -1,113 +0,0 @@
|
||||
|
||||
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
|
||||
|
||||
## Video Sparse Attention (VSA)
|
||||
|
||||
### Installation
|
||||
We support H100 (via TK) and any other GPU (via triton) for VSA.
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
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
|
||||
```
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
### Verify if you have successfully installed
|
||||
|
||||
```bash
|
||||
# test numerical
|
||||
python tests/test_vsa.py
|
||||
# (For H100) test speed
|
||||
python benchmarks/bench_vsa_hopper.py
|
||||
```
|
||||
bench_vsa_hopper.py should print something like this:
|
||||
```bash
|
||||
Using topk=76 kv blocks per q block (out of 768 total kv blocks)
|
||||
|
||||
=== BLOCK SPARSE ATTENTION BENCHMARK ===
|
||||
Block Sparse Forward - TFLOPS: 5622.26
|
||||
Block Sparse Backward - TFLOPS: 3865.68
|
||||
```
|
||||
|
||||
|
||||
## Sliding Tile Attention (STA)
|
||||
We only support H100 for STA.
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup_sta.py install
|
||||
```
|
||||
|
||||
|
||||
|
||||
|
||||
### Usage
|
||||
End-2-end inference with FastVideo:
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_STA.sh
|
||||
```
|
||||
|
||||
If you want to use sliding tile attention in your custom model:
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
```
|
||||
|
||||
|
||||
### Test
|
||||
```bash
|
||||
python tests/test_sta.py # test STA
|
||||
python tests/test_vsa.py # test VSA
|
||||
```
|
||||
### Benchmark
|
||||
```bash
|
||||
python benchmarks/bench_sta.py
|
||||
```
|
||||
|
||||
|
||||
### How Does STA Work?
|
||||
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
|
||||
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
|
||||
|
||||
## Why is STA Fast?
|
||||
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
|
||||
|
||||
STA removes mixed blocks.
|
||||
|
||||
|
||||
<div align="center">
|
||||
<img src=../../assets/sliding_tile_attn_map.png width="80%"/>
|
||||
</div>
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
|
||||
@@ -1,224 +0,0 @@
|
||||
import torch
|
||||
import argparse
|
||||
from triton.testing import do_bench
|
||||
from vsa import block_sparse_fwd, block_sparse_bwd
|
||||
from vsa import BLOCK_M, BLOCK_N
|
||||
import triton
|
||||
import numpy as np
|
||||
import random
|
||||
|
||||
def set_seed(seed: int = 42):
|
||||
# Python random module
|
||||
random.seed(seed)
|
||||
|
||||
# NumPy
|
||||
np.random.seed(seed)
|
||||
|
||||
# PyTorch
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed) # if using multi-GPU
|
||||
|
||||
def parse_arguments():
|
||||
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
|
||||
parser.add_argument('--batch_size', type=int, default=1, help='Batch size')
|
||||
parser.add_argument('--num_heads', type=int, default=12, help='Number of heads')
|
||||
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=None, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[49152], help='Sequence lengths to benchmark')
|
||||
return parser.parse_args()
|
||||
|
||||
def create_input_tensors(batch, head, seq_len, headdim):
|
||||
"""Create random input tensors for attention."""
|
||||
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
return q, k, v
|
||||
|
||||
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
|
||||
|
||||
Args:
|
||||
bs: batch size
|
||||
h: number of heads
|
||||
num_q_blocks: number of query blocks
|
||||
num_kv_blocks: number of key-value blocks
|
||||
k: number of kv blocks each q block attends to
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to k).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Binary mask where 1 indicates attention connection.
|
||||
"""
|
||||
# Ensure k is not larger than num_kv_blocks
|
||||
k = min(k, num_kv_blocks)
|
||||
|
||||
# Create random scores for sampling
|
||||
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
|
||||
|
||||
# Get top-k indices for each q block
|
||||
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
|
||||
|
||||
# sort q2k_block_sparse_index
|
||||
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
|
||||
|
||||
# All q blocks attend to exactly k kv blocks
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
|
||||
|
||||
# Create the corresponding mask
|
||||
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
|
||||
|
||||
# Fill in the mask based on the indices
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx]
|
||||
block_sparse_mask[b, head, q_idx, kv_indices] = True
|
||||
|
||||
# Create the reverse mapping (k2q)
|
||||
# First, initialize lists to collect q indices for each kv block
|
||||
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
|
||||
|
||||
# Populate the lists based on q2k mapping
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
|
||||
for kv_idx in kv_indices:
|
||||
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
|
||||
|
||||
# Find the maximum number of q blocks that attend to any kv block
|
||||
max_q_per_kv = 0
|
||||
for flat_idx in range(bs * h):
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
|
||||
|
||||
# Create tensors for k2q mapping
|
||||
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
|
||||
dtype=torch.int32, device=device)
|
||||
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
|
||||
dtype=torch.int32, device=device)
|
||||
|
||||
# Fill the tensors
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
q_indices = k2q_indices_list[flat_idx][kv_idx]
|
||||
num_q = len(q_indices)
|
||||
k2q_block_sparse_num[b, head, kv_idx] = num_q
|
||||
if num_q > 0:
|
||||
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
|
||||
q_indices, dtype=torch.int32, device=device)
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
|
||||
|
||||
def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops):
|
||||
"""Benchmark block sparse attention forward and backward passes."""
|
||||
print("\n=== BLOCK SPARSE ATTENTION BENCHMARK ===")
|
||||
|
||||
# Forward pass
|
||||
# Warm-up run
|
||||
variable_block_sizes = torch.ones(q2k_block_sparse_index.shape[2], device=q.device).int() * BLOCK_M
|
||||
o, l_vec = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Benchmark forward
|
||||
fwd_time = do_bench(
|
||||
lambda: block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes),
|
||||
warmup=5,
|
||||
rep=20,
|
||||
quantiles=None
|
||||
)
|
||||
|
||||
sparse_tflops = flops / fwd_time * 1e-12 * 1e3
|
||||
print(f"Block Sparse Forward - TFLOPS: {sparse_tflops:.2f}")
|
||||
|
||||
# Backward pass
|
||||
grad_output = torch.randn_like(o)
|
||||
|
||||
# Warm-up runs
|
||||
for _ in range(5):
|
||||
block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Benchmark backward
|
||||
bwd_time = do_bench(
|
||||
lambda: block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes),
|
||||
warmup=5,
|
||||
rep=20,
|
||||
quantiles=None
|
||||
)
|
||||
bwd_flops = 2.5 * flops # Approximation
|
||||
|
||||
sparse_bwd_tflops = bwd_flops / bwd_time * 1e-12 * 1e3
|
||||
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd_tflops:.2f}")
|
||||
|
||||
return sparse_tflops, sparse_bwd_tflops
|
||||
|
||||
def main():
|
||||
args = parse_arguments()
|
||||
|
||||
set_seed(42)
|
||||
|
||||
# Extract parameters
|
||||
batch = args.batch_size
|
||||
head = args.num_heads
|
||||
headdim = args.head_dim
|
||||
|
||||
print(f"Block Sparse Attention Benchmark")
|
||||
print(f"batch: {batch}, head: {head}, headdim: {headdim}")
|
||||
|
||||
# Test with different sequence lengths
|
||||
for seq_len in args.seq_lengths:
|
||||
# Skip very long sequences if they might cause OOM
|
||||
if seq_len > 16384 and batch > 1:
|
||||
continue
|
||||
|
||||
print("="*100)
|
||||
print(f"\nSequence length: {seq_len}")
|
||||
|
||||
# Calculate theoretical FLOPs for attention
|
||||
flops = 4 * batch * head * headdim * seq_len * seq_len
|
||||
|
||||
# Create input tensors
|
||||
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
|
||||
|
||||
# Setup block sparse parameters
|
||||
num_q_blocks = seq_len // BLOCK_M
|
||||
num_kv_blocks = seq_len // BLOCK_N
|
||||
|
||||
# Determine k value (number of kv blocks per q block)
|
||||
topk = args.topk
|
||||
if topk is None:
|
||||
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
|
||||
topk = max(1, topk)
|
||||
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
|
||||
|
||||
# Generate block sparse pattern
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, _ = generate_block_sparse_pattern(
|
||||
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
|
||||
|
||||
# Benchmark block sparse attention
|
||||
sparse_fwd, sparse_bwd = benchmark_block_sparse_attention(
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops
|
||||
)
|
||||
|
||||
# Print results
|
||||
print("\n=== PERFORMANCE RESULTS ===")
|
||||
print(f"Block Sparse Forward - TFLOPS: {sparse_fwd:.2f}")
|
||||
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd:.2f}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,217 +0,0 @@
|
||||
import torch
|
||||
import argparse
|
||||
import triton.testing
|
||||
from vsa import block_sparse_attn
|
||||
from vsa import BLOCK_M, BLOCK_N
|
||||
|
||||
import numpy as np
|
||||
import random
|
||||
|
||||
def set_seed(seed: int = 42):
|
||||
# Python random module
|
||||
random.seed(seed)
|
||||
|
||||
# NumPy
|
||||
np.random.seed(seed)
|
||||
|
||||
# PyTorch
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed) # if using multi-GPU
|
||||
|
||||
def parse_arguments():
|
||||
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
|
||||
parser.add_argument('--batch_size', type=int, default=1, help='Batch size')
|
||||
parser.add_argument('--num_heads', type=int, default=12, help='Number of heads')
|
||||
parser.add_argument('--head_dim', type=int, default=64, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=None, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[49152], help='Sequence lengths to benchmark')
|
||||
return parser.parse_args()
|
||||
|
||||
def create_input_tensors(batch, head, seq_len, headdim):
|
||||
"""Create random input tensors for attention."""
|
||||
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
return q, k, v
|
||||
|
||||
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
|
||||
|
||||
Args:
|
||||
bs: batch size
|
||||
h: number of heads
|
||||
num_q_blocks: number of query blocks
|
||||
num_kv_blocks: number of key-value blocks
|
||||
k: number of kv blocks each q block attends to
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to k).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Binary mask where 1 indicates attention connection.
|
||||
"""
|
||||
# Ensure k is not larger than num_kv_blocks
|
||||
k = min(k, num_kv_blocks)
|
||||
|
||||
# Create random scores for sampling
|
||||
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
|
||||
|
||||
# Get top-k indices for each q block
|
||||
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
|
||||
|
||||
# sort q2k_block_sparse_index
|
||||
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
|
||||
|
||||
# All q blocks attend to exactly k kv blocks
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
|
||||
|
||||
# Create the corresponding mask
|
||||
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
|
||||
|
||||
# Fill in the mask based on the indices
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx]
|
||||
block_sparse_mask[b, head, q_idx, kv_indices] = True
|
||||
|
||||
# Create the reverse mapping (k2q)
|
||||
# First, initialize lists to collect q indices for each kv block
|
||||
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
|
||||
|
||||
# Populate the lists based on q2k mapping
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
|
||||
for kv_idx in kv_indices:
|
||||
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
|
||||
|
||||
# Find the maximum number of q blocks that attend to any kv block
|
||||
max_q_per_kv = 0
|
||||
for flat_idx in range(bs * h):
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
|
||||
|
||||
# Create tensors for k2q mapping
|
||||
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
|
||||
dtype=torch.int32, device=device)
|
||||
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
|
||||
dtype=torch.int32, device=device)
|
||||
|
||||
# Fill the tensors
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
q_indices = k2q_indices_list[flat_idx][kv_idx]
|
||||
num_q = len(q_indices)
|
||||
k2q_block_sparse_num[b, head, kv_idx] = num_q
|
||||
if num_q > 0:
|
||||
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
|
||||
q_indices, dtype=torch.int32, device=device)
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
|
||||
|
||||
def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops):
|
||||
"""Benchmark block sparse attention forward+backward pass."""
|
||||
print("\n=== BLOCK SPARSE ATTENTION FORWARD+BACKWARD BENCHMARK ===")
|
||||
|
||||
# Combined forward+backward pass
|
||||
# Warm-up run
|
||||
q_fwd = q.clone().requires_grad_(True)
|
||||
k_fwd = k.clone().requires_grad_(True)
|
||||
v_fwd = v.clone().requires_grad_(True)
|
||||
o = block_sparse_attn(q_fwd, k_fwd, v_fwd, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
grad_output = torch.randn_like(o)
|
||||
o.backward(grad_output)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Benchmark forward+backward
|
||||
def forward_backward_fn():
|
||||
q_fwd = q.clone().requires_grad_(True)
|
||||
k_fwd = k.clone().requires_grad_(True)
|
||||
v_fwd = v.clone().requires_grad_(True)
|
||||
o = block_sparse_attn(q_fwd, k_fwd, v_fwd, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
grad_output = torch.randn_like(o)
|
||||
o.backward(grad_output)
|
||||
|
||||
total_time = triton.testing.do_bench(
|
||||
forward_backward_fn,
|
||||
warmup=25,
|
||||
rep=100,
|
||||
return_mode='mean'
|
||||
)
|
||||
|
||||
# Total flops for forward + backward (forward + 2.5x backward approximation)
|
||||
total_flops = flops + 2.5 * flops # 3.5x the forward flops
|
||||
sparse_tflops = total_flops / total_time * 1e-12 * 1e3
|
||||
print(f"Block Sparse Forward+Backward - TFLOPS: {sparse_tflops:.2f}")
|
||||
|
||||
return sparse_tflops
|
||||
|
||||
def main():
|
||||
args = parse_arguments()
|
||||
|
||||
set_seed(42)
|
||||
|
||||
# Extract parameters
|
||||
batch = args.batch_size
|
||||
head = args.num_heads
|
||||
headdim = args.head_dim
|
||||
|
||||
print(f"Block Sparse Attention Benchmark")
|
||||
print(f"batch: {batch}, head: {head}, headdim: {headdim}")
|
||||
|
||||
# Test with different sequence lengths
|
||||
for seq_len in args.seq_lengths:
|
||||
# Skip very long sequences if they might cause OOM
|
||||
if seq_len > 16384 and batch > 1:
|
||||
continue
|
||||
|
||||
print("="*100)
|
||||
print(f"\nSequence length: {seq_len}")
|
||||
|
||||
# Calculate theoretical FLOPs for attention
|
||||
flops = 4 * batch * head * headdim * seq_len * seq_len
|
||||
|
||||
# Create input tensors
|
||||
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
|
||||
|
||||
# Setup block sparse parameters
|
||||
num_q_blocks = seq_len // BLOCK_M
|
||||
num_kv_blocks = seq_len // BLOCK_N
|
||||
|
||||
# Determine k value (number of kv blocks per q block)
|
||||
topk = args.topk
|
||||
if topk is None:
|
||||
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
|
||||
topk = max(1, topk)
|
||||
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
|
||||
|
||||
# Generate block sparse pattern
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, _ = generate_block_sparse_pattern(
|
||||
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
|
||||
|
||||
# Benchmark block sparse attention
|
||||
sparse_fwd = benchmark_block_sparse_attention(
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops
|
||||
)
|
||||
|
||||
# Print results
|
||||
print("\n=== PERFORMANCE RESULTS ===")
|
||||
print(f"Block Sparse Forward+Backward - TFLOPS: {sparse_fwd:.2f}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,156 +0,0 @@
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
# Add the parent directory to the path to import block_sparse_attn
|
||||
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
from tests.utils import generate_block_sparse_mask_for_function, create_full_mask_from_block_mask
|
||||
from vsa import block_sparse_attn
|
||||
|
||||
BLOCK_M = 64
|
||||
BLOCK_N = 64
|
||||
|
||||
def pytorch_test(Q, K, V, block_sparse_mask, dO):
|
||||
q_ = Q.clone().float().requires_grad_()
|
||||
k_ = K.clone().float().requires_grad_()
|
||||
v_ = V.clone().float().requires_grad_()
|
||||
|
||||
QK = torch.matmul(q_, k_.transpose(-2, -1))
|
||||
QK /= (q_.size(-1) ** 0.5)
|
||||
QK = QK.masked_fill(~block_sparse_mask.unsqueeze(0), float('-inf'))
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v_)
|
||||
|
||||
dO_ = dO
|
||||
output.backward(dO_)
|
||||
return (
|
||||
output.to(torch.bfloat16),
|
||||
q_.grad.to(torch.bfloat16),
|
||||
k_.grad.to(torch.bfloat16),
|
||||
v_.grad.to(torch.bfloat16),
|
||||
)
|
||||
|
||||
|
||||
def block_sparse_kernel_test(Q, K, V, block_sparse_mask, variable_block_sizes, non_pad_index, dO):
|
||||
Q = Q.detach().requires_grad_()
|
||||
K = K.detach().requires_grad_()
|
||||
V = V.detach().requires_grad_()
|
||||
|
||||
q_padded = vsa_pad(Q, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
|
||||
k_padded = vsa_pad(K, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
|
||||
v_padded = vsa_pad(V, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
|
||||
output, _= block_sparse_attn(q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes)
|
||||
output = output[:, :, non_pad_index, :]
|
||||
output.backward(dO)
|
||||
return output, Q.grad, K.grad, V.grad
|
||||
|
||||
|
||||
def get_non_pad_index(
|
||||
vid_len: torch.LongTensor,
|
||||
n_win: int,
|
||||
win_size: int,
|
||||
):
|
||||
device = vid_len.device
|
||||
starts_pad = torch.arange(n_win, device=device) * win_size
|
||||
index_pad = starts_pad[:, None] + torch.arange(win_size, device=device)[None, :]
|
||||
index_mask = torch.arange(win_size, device=device)[None, :] < vid_len[:, None]
|
||||
|
||||
return index_pad[index_mask]
|
||||
|
||||
def generate_tensor(shape, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
return tensor
|
||||
|
||||
def generate_variable_block_sizes(num_blocks, min_size=32, max_size=64, device="cuda"):
|
||||
return torch.randint(min_size, max_size + 1, (num_blocks,), device=device, dtype=torch.int32)
|
||||
|
||||
|
||||
def vsa_pad(x, non_pad_index, num_blocks, block_size):
|
||||
padded_x = torch.zeros((1, x.shape[1], num_blocks * BLOCK_M, x.shape[3]), device=x.device, dtype=x.dtype)
|
||||
padded_x[:, :, non_pad_index, :] = x
|
||||
return padded_x
|
||||
|
||||
def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all'):
|
||||
results = {
|
||||
'gO': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
'gQ': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
'gK': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
'gV': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
}
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
variable_block_sizes = generate_variable_block_sizes(num_blocks, device=device)
|
||||
S = int(variable_block_sizes.sum().item())
|
||||
padded_S = num_blocks * BLOCK_M
|
||||
non_pad_index = get_non_pad_index(variable_block_sizes, num_blocks, BLOCK_M)
|
||||
block_mask = generate_block_sparse_mask_for_function(h, num_blocks, k, device)
|
||||
full_mask = create_full_mask_from_block_mask(block_mask, variable_block_sizes, device)
|
||||
for _ in range(num_iterations):
|
||||
Q = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
K = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
V = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
dO = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
|
||||
# dO_padded = torch.zeros_like(dO_padded)
|
||||
# dO_padded[:, :, non_pad_index, :] = dO
|
||||
|
||||
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, full_mask, dO)
|
||||
bs_o, bs_qg, bs_kg, bs_vg = block_sparse_kernel_test(Q, K, V, block_mask.unsqueeze(0), variable_block_sizes,non_pad_index, dO)
|
||||
for name, (pt, bs) in zip(['gQ', 'gK', 'gV', 'gO'], [(pt_qg, bs_qg), (pt_kg, bs_kg), (pt_vg, bs_vg), (pt_o, bs_o)]):
|
||||
if bs is not None:
|
||||
diff = pt - bs
|
||||
abs_diff = torch.abs(diff)
|
||||
results[name]['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results[name]['sum_abs'] += torch.sum(torch.abs(pt)).item()
|
||||
rel_max_diff = torch.max(abs_diff) / torch.mean(torch.abs(pt))
|
||||
results[name]['max_diff'] = max(results[name]['max_diff'], rel_max_diff.item())
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
total_elements = h * S * d * num_iterations
|
||||
for name, data in results.items():
|
||||
avg_diff = data['sum_diff'] / total_elements
|
||||
max_diff = data['max_diff']
|
||||
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
|
||||
|
||||
return results
|
||||
|
||||
def generate_error_graphs(h, d, error_mode='all'):
|
||||
test_configs = [
|
||||
{"num_blocks": 16, "k": 2, "description": "Small sequence"},
|
||||
{"num_blocks": 32, "k": 4, "description": "Medium sequence"},
|
||||
{"num_blocks": 53, "k": 6, "description": "Large sequence"},
|
||||
]
|
||||
|
||||
print(f"\nError Analysis for h={h}, d={d}, mode={error_mode}")
|
||||
print("=" * 150)
|
||||
print(f"{'Config':<20} {'Blocks':<8} {'K':<4} "
|
||||
f"{'gQ Avg':<12} {'Rel gQ Max':<12} "
|
||||
f"{'gK Avg':<12} {'Rel gK Max':<12} "
|
||||
f"{'gV Avg':<12} {'Rel gV Max':<12} "
|
||||
f"{'gO Avg':<12} {'Rel gO Max':<12}")
|
||||
print("-" * 150)
|
||||
|
||||
for config in test_configs:
|
||||
num_blocks = config["num_blocks"]
|
||||
k = config["k"]
|
||||
description = config["description"]
|
||||
results = check_correctness(h, d, num_blocks, k, error_mode=error_mode)
|
||||
print(f"{description:<20} {num_blocks:<8} {k:<4} "
|
||||
f"{results['gQ']['avg_diff']:<12.6e} {results['gQ']['max_diff']:<12.6e} "
|
||||
f"{results['gK']['avg_diff']:<12.6e} {results['gK']['max_diff']:<12.6e} "
|
||||
f"{results['gV']['avg_diff']:<12.6e} {results['gV']['max_diff']:<12.6e} "
|
||||
f"{results['gO']['avg_diff']:<12.6e} {results['gO']['max_diff']:<12.6e}")
|
||||
|
||||
print("-" * 150)
|
||||
|
||||
if __name__ == "__main__":
|
||||
h, d = 16, 128
|
||||
print("Block Sparse Attention with Variable Block Sizes Analysis")
|
||||
print("=" * 60)
|
||||
for mode in ['backward']:
|
||||
generate_error_graphs(h, d, error_mode=mode)
|
||||
print("\nAnalysis completed for all modes.")
|
||||
@@ -1,54 +0,0 @@
|
||||
import torch
|
||||
|
||||
def generate_block_sparse_mask_for_function(h, num_blocks, k, device="cuda"):
|
||||
"""
|
||||
Generate block sparse mask of shape [h, num_blocks, num_blocks].
|
||||
|
||||
Args:
|
||||
h: number of heads
|
||||
num_blocks: number of blocks
|
||||
k: number of kv blocks each q block attends to
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
|
||||
"""
|
||||
k = min(k, num_blocks)
|
||||
scores = torch.rand(h, num_blocks, num_blocks, device=device)
|
||||
_, indices = torch.topk(scores, k, dim=-1)
|
||||
block_sparse_mask = torch.zeros(h, num_blocks, num_blocks, dtype=torch.bool, device=device)
|
||||
|
||||
block_sparse_mask = block_sparse_mask.scatter_(2, indices, 1).bool()
|
||||
return block_sparse_mask
|
||||
|
||||
|
||||
def create_full_mask_from_block_mask(block_sparse_mask, variable_block_sizes, device="cuda"):
|
||||
"""
|
||||
Convert block-level sparse mask to full attention mask.
|
||||
|
||||
Args:
|
||||
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
|
||||
variable_block_sizes: [num_blocks] tensor
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
full_mask: [h, S, S] bool tensor where S = total sequence length
|
||||
"""
|
||||
h, num_blocks, _ = block_sparse_mask.shape
|
||||
total_seq_len = variable_block_sizes.sum().item()
|
||||
cumsum = torch.cat([torch.tensor([0], device=device), variable_block_sizes.cumsum(dim=0)[:-1]])
|
||||
|
||||
full_mask = torch.zeros(h, total_seq_len, total_seq_len, dtype=torch.bool, device=device)
|
||||
|
||||
for head in range(h):
|
||||
for q_block in range(num_blocks):
|
||||
q_start = cumsum[q_block]
|
||||
q_end = q_start + variable_block_sizes[q_block]
|
||||
|
||||
for kv_block in range(num_blocks):
|
||||
if block_sparse_mask[head, q_block, kv_block]:
|
||||
kv_start = cumsum[kv_block]
|
||||
kv_end = kv_start + variable_block_sizes[kv_block]
|
||||
full_mask[head, q_start:q_end, kv_start:kv_end] = True
|
||||
|
||||
return full_mask
|
||||
@@ -1,2 +0,0 @@
|
||||
recursive-include tk *
|
||||
include config_vsa.py
|
||||
@@ -1,61 +0,0 @@
|
||||
|
||||
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Video Sparse Attention (VSA)
|
||||
|
||||
### Installation
|
||||
We support H100 (via TK) and any other GPU (via triton) for VSA.
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
Install from source:
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
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
|
||||
```
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
### Verify if you have successfully installed
|
||||
|
||||
```bash
|
||||
# test numerical
|
||||
python ../tests/test_vsa.py
|
||||
# (For H100) test speed
|
||||
python ../benchmarks/bench_vsa_hopper.py
|
||||
```
|
||||
|
||||
bench_vsa_hopper.py should print something like this:
|
||||
|
||||
```bash
|
||||
Using topk=76 kv blocks per q block (out of 768 total kv blocks)
|
||||
|
||||
=== BLOCK SPARSE ATTENTION BENCHMARK ===
|
||||
Block Sparse Forward - TFLOPS: 5622.26
|
||||
Block Sparse Backward - TFLOPS: 3865.68
|
||||
```
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
|
||||
@@ -1,15 +0,0 @@
|
||||
### ADD TO THIS TO REGISTER NEW KERNELS
|
||||
sources = {
|
||||
'block_sparse': {
|
||||
'source_files': {
|
||||
'h100': 'vsa/block_sparse_h100.cu'
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
### WHICH KERNELS DO WE WANT TO BUILD?
|
||||
# (oftentimes during development work you don't need to redefine them all.)
|
||||
kernels = ['block_sparse']
|
||||
|
||||
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
|
||||
target = 'h100'
|
||||
@@ -1,81 +0,0 @@
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
from config_vsa import kernels, sources, target
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "vsa"
|
||||
VERSION = "0.0.3"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Video Sparse Attention Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn/video_sparse_attn"
|
||||
|
||||
# Set environment variables
|
||||
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
|
||||
python_include = subprocess.check_output(['python', '-c',
|
||||
"import sysconfig; print(sysconfig.get_path('include'))"]).decode().strip()
|
||||
torch_include = subprocess.check_output([
|
||||
'python', '-c',
|
||||
"import torch; from torch.utils.cpp_extension import include_paths; print(' '.join(['-I' + p for p in include_paths()]))"
|
||||
]).decode().strip()
|
||||
print('vsa root:', tk_root)
|
||||
print('Python include:', python_include)
|
||||
print('Torch include directories:', torch_include)
|
||||
|
||||
# CUDA flags
|
||||
cuda_flags = [
|
||||
'-DNDEBUG', '-Xcompiler=-Wno-psabi', '-Xcompiler=-fno-strict-aliasing', '--expt-extended-lambda',
|
||||
'--expt-relaxed-constexpr', '-forward-unknown-to-host-compiler', '--use_fast_math', '-std=c++20', '-O3',
|
||||
'-Xnvlink=--verbose', '-Xptxas=--verbose', '-Xptxas=--warn-on-spills', f'-I{tk_root}/include',
|
||||
f'-I{tk_root}/prototype', f'-I{python_include}', '-DTORCH_COMPILE'
|
||||
] + torch_include.split()
|
||||
cpp_flags = ['-std=c++20', '-O3']
|
||||
|
||||
if target == 'h100':
|
||||
cuda_flags.append('-DKITTENS_HOPPER')
|
||||
cuda_flags.append('-arch=sm_90a')
|
||||
else:
|
||||
raise ValueError(f'Target {target} not supported')
|
||||
|
||||
source_files = ['vsa.cpp']
|
||||
for k in kernels:
|
||||
if target not in sources[k]['source_files']:
|
||||
raise KeyError(f'Target {target} not found in source files for kernel {k}')
|
||||
if isinstance(sources[k]['source_files'][target], list):
|
||||
source_files.extend(sources[k]['source_files'][target])
|
||||
else:
|
||||
source_files.append(sources[k]['source_files'][target])
|
||||
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
|
||||
|
||||
|
||||
ext_modules = [
|
||||
CUDAExtension('vsa_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
]
|
||||
|
||||
|
||||
|
||||
setup(name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
ext_modules=ext_modules,
|
||||
cmdclass={'build_ext': BuildExtension},
|
||||
classifiers=[
|
||||
"Programming Language :: Python :: 3",
|
||||
"Environment :: GPU :: NVIDIA CUDA :: 12",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
],
|
||||
python_requires='>=3.10',
|
||||
install_requires=["torch>=2.5.0"])
|
||||
@@ -1,27 +0,0 @@
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/ATen.h>
|
||||
|
||||
#include <vector>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE
|
||||
extern std::vector<torch::Tensor> block_sparse_attention_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num, torch::Tensor block_size
|
||||
);
|
||||
extern std::vector<torch::Tensor> block_sparse_attention_backward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, torch::Tensor l_vec, torch::Tensor og, torch::Tensor k2q_block_sparse_index, torch::Tensor k2q_block_sparse_num, torch::Tensor block_size
|
||||
);
|
||||
#endif
|
||||
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "Video Sparse Attention Kernels"; // optional module docstring
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE
|
||||
m.def("block_sparse_fwd", torch::wrap_pybind_function(block_sparse_attention_forward), "block sparse attention");
|
||||
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward");
|
||||
#endif
|
||||
}
|
||||
@@ -1,80 +0,0 @@
|
||||
import torch
|
||||
from typing import Tuple
|
||||
block_sparse_attn=None
|
||||
import torch
|
||||
major, minor = torch.cuda.get_device_capability(0)
|
||||
if major == 9 and minor == 0:# check if H100
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
from vsa.block_sparse_wrapper import block_sparse_attn_SM90
|
||||
block_sparse_attn = block_sparse_attn_SM90
|
||||
else:
|
||||
from vsa.block_sparse_wrapper import block_sparse_attn_triton
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
block_sparse_attn = block_sparse_attn_triton
|
||||
|
||||
BLOCK_M = 64
|
||||
BLOCK_N = 64
|
||||
|
||||
|
||||
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
QK = torch.matmul(q, k.transpose(-2, -1))
|
||||
QK /= (q.size(-1)**0.5)
|
||||
|
||||
# Causal mask removed since causal is always false
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v)
|
||||
return output, QK
|
||||
|
||||
|
||||
def video_sparse_attn(q, k, v, variable_block_sizes, topk, block_size, compress_attn_weight=None):
|
||||
"""
|
||||
q: [batch_size, num_heads, seq_len, head_dim]
|
||||
k: [batch_size, num_heads, seq_len, head_dim]
|
||||
v: [batch_size, num_heads, seq_len, head_dim]
|
||||
topk: int
|
||||
block_size: int or tuple of 3 ints
|
||||
video_shape: tuple of (T, H, W)
|
||||
compress_attn_weight: [batch_size, num_heads, seq_len, head_dim]
|
||||
select_attn_weight: [batch_size, num_heads, seq_len, head_dim]
|
||||
NOTE: We assume q, k, v is zero padded!!
|
||||
V1 of sparse attention. Include compress attn and sparse attn branch, use average pooling to compress.
|
||||
Assume q, k, v is flattened in this way: [batch_size, num_heads, T//block_size[0], H//block_size[1], W//block_size[2], block_size[0], block_size[1], block_size[2]]
|
||||
"""
|
||||
|
||||
if isinstance(block_size, int):
|
||||
block_size = (block_size, block_size, block_size)
|
||||
|
||||
block_elements = block_size[0] * block_size[1] * block_size[2]
|
||||
assert block_elements == 64
|
||||
assert q.shape[2] % block_elements == 0
|
||||
batch_size, num_heads, seq_len, head_dim = q.shape
|
||||
# compress attn
|
||||
q_compress = (q.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(q.dtype)
|
||||
k_compress = (k.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(k.dtype)
|
||||
v_compress = (v.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(v.dtype)
|
||||
|
||||
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
|
||||
v_compress)
|
||||
|
||||
output_compress = output_compress.view(batch_size, num_heads,
|
||||
seq_len // block_elements, 1,
|
||||
head_dim)
|
||||
output_compress = output_compress.repeat(1, 1, 1, block_elements,
|
||||
1).view(batch_size, num_heads,
|
||||
seq_len, head_dim)
|
||||
|
||||
topK_indices = torch.topk(block_attn_score, topk, dim=-1).indices
|
||||
block_mask = torch.zeros_like(block_attn_score, dtype=torch.bool).scatter_(-1, topK_indices, True)
|
||||
output_select, _ = block_sparse_attn(q, k, v, block_mask, variable_block_sizes)
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
final_output = output_compress * compress_attn_weight + output_select
|
||||
else:
|
||||
final_output = output_compress + output_select
|
||||
return final_output
|
||||
|
||||
@@ -1,449 +0,0 @@
|
||||
"""
|
||||
Fused Attention
|
||||
===============
|
||||
|
||||
This is a Triton implementation of the Flash Attention v2 algorithm from Tri Dao
|
||||
(https://tridao.me/publications/flash2/flash2.pdf)
|
||||
|
||||
Credits: OpenAI kernel team
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
import math # small utility needed by the sparse wrapper
|
||||
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
|
||||
|
||||
|
||||
# We don't run auto-tuning every time to keep the tutorial fast. Keeping
|
||||
# the code below and commenting out the equivalent parameters is convenient for
|
||||
# re-tuning.
|
||||
configs = [
|
||||
triton.Config({'BLOCK_M': BM, 'BLOCK_N': BN}, num_stages=s, num_warps=w) \
|
||||
for BM in [64]\
|
||||
for BN in [64]\
|
||||
for s in [3, 4, 7]\
|
||||
for w in [4, 8]\
|
||||
]
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
@triton.autotune(configs, key=["N_CTX", "HEAD_DIM"])
|
||||
@triton.jit
|
||||
def _attn_fwd_sparse(Q, K, V, sm_scale, #
|
||||
q2k_index, q2k_num, max_kv_blks, #
|
||||
variable_block_sizes,
|
||||
M, Out, #
|
||||
stride_qz, stride_qh, stride_qm, stride_qk,
|
||||
stride_kz, stride_kh, stride_kn, stride_kk,
|
||||
stride_vz, stride_vh, stride_vk, stride_vn,
|
||||
stride_oz, stride_oh, stride_om, stride_on,
|
||||
Z, H, N_CTX, #
|
||||
HEAD_DIM: tl.constexpr, #
|
||||
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
|
||||
STAGE: tl.constexpr):
|
||||
"""
|
||||
64×64 **block-sparse** forward kernel. Back-prop kernels remain dense
|
||||
(32×64 and 64×32) – memory footprint unchanged.
|
||||
"""
|
||||
|
||||
# ----- program-id mapping -----
|
||||
q_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(1) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_M
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
|
||||
# ----- base pointers -----
|
||||
qvk_off = (b.to(tl.int64) * stride_qz +
|
||||
h.to(tl.int64) * stride_qh)
|
||||
|
||||
Q_ptr = tl.make_block_ptr(
|
||||
base=Q + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_qm, stride_qk),
|
||||
offsets=(q_blk * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
|
||||
|
||||
K_base = tl.make_block_ptr(
|
||||
base=K + qvk_off, shape=(HEAD_DIM, N_CTX),
|
||||
strides=(stride_kk, stride_kn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_N), order=(0, 1))
|
||||
|
||||
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1, 0)
|
||||
V_base = tl.make_block_ptr(
|
||||
base=V + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_vk, stride_vn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(BLOCK_N, HEAD_DIM), order=v_order)
|
||||
|
||||
O_ptr = tl.make_block_ptr(
|
||||
base=Out + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_om, stride_on),
|
||||
offsets=(q_blk * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
|
||||
|
||||
# ----- accumulators -----
|
||||
offs_m = q_blk * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
m_i = tl.full([BLOCK_M], -float("inf"), tl.float32)
|
||||
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
|
||||
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
|
||||
qk_scale = sm_scale * 1.44269504 # 1/ln2
|
||||
q = tl.load(Q_ptr)
|
||||
|
||||
# ----- sparse loop over valid K/V tiles -----
|
||||
for i in range(0, kv_blocks):
|
||||
kv_idx = tl.load(kv_ptr + i).to(tl.int32)
|
||||
block_size = tl.load(variable_block_sizes + kv_idx)
|
||||
K_ptr = tl.advance(K_base, (0, kv_idx * BLOCK_N))
|
||||
V_ptr = tl.advance(V_base, (kv_idx * BLOCK_N, 0))
|
||||
|
||||
k = tl.load(K_ptr)
|
||||
qk = tl.dot(q, k)
|
||||
# mask out invalid columns
|
||||
mask = tl.arange(0, BLOCK_N) < block_size
|
||||
qk = tl.where(mask[None, :], qk, -float("inf"))
|
||||
|
||||
m_ij = tl.maximum(m_i, tl.max(qk, 1) * qk_scale)
|
||||
p = tl.math.exp2(qk * qk_scale - m_ij[:, None])
|
||||
l_ij = tl.sum(p, 1)
|
||||
|
||||
alpha = tl.math.exp2(m_i - m_ij)
|
||||
l_i = l_i * alpha + l_ij
|
||||
acc = acc * alpha[:, None]
|
||||
|
||||
v = tl.load(V_ptr)
|
||||
acc = tl.dot(p.to(tl.bfloat16), v, acc)
|
||||
m_i = m_ij
|
||||
|
||||
# ----- epilogue -----
|
||||
m_i += tl.math.log2(l_i)
|
||||
acc = acc / l_i[:, None]
|
||||
tl.store(M + off_hz * N_CTX + offs_m, m_i)
|
||||
tl.store(O_ptr, acc.to(Out.type.element_ty))
|
||||
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
|
||||
|
||||
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd_preprocess(O, DO, #
|
||||
Delta, #
|
||||
Z, H, N_CTX, #
|
||||
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr #
|
||||
):
|
||||
off_m = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
off_hz = tl.program_id(1)
|
||||
off_n = tl.arange(0, HEAD_DIM)
|
||||
# load
|
||||
o = tl.load(O + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :])
|
||||
do = tl.load(DO + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :]).to(tl.float32)
|
||||
delta = tl.sum(o * do, axis=1)
|
||||
# write-back
|
||||
tl.store(Delta + off_hz * N_CTX + off_m, delta)
|
||||
|
||||
|
||||
# The main inner-loop logic for computing dK and dV.
|
||||
@triton.jit
|
||||
def _attn_bwd_dkdv(dk, dv, #
|
||||
Q, k, v, sm_scale, #
|
||||
DO, #
|
||||
M, D, #
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
# shared by Q/K/V/DO.
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, BLOCK_M1: tl.constexpr, #
|
||||
BLOCK_N1: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr, #
|
||||
# Filled in by the wrapper.
|
||||
start_n, start_m, num_steps):
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M1)
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N1)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
qT_ptrs = Q + offs_m[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
do_ptrs = DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
# BLOCK_N1 must be a multiple of BLOCK_M1, otherwise the code wouldn't work.
|
||||
tl.static_assert(BLOCK_N1 % BLOCK_M1 == 0)
|
||||
step_m = BLOCK_M1
|
||||
kv_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(2) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_N1
|
||||
meta_base = ((b * H + h) * q_tiles + kv_blk)
|
||||
|
||||
q_blocks = tl.load(k2q_num + meta_base) # int32
|
||||
q_ptr = k2q_index + meta_base * max_q_blks # ptr to list
|
||||
block_size = tl.load(variable_block_sizes + kv_blk)
|
||||
|
||||
|
||||
|
||||
for blk_idx in range(q_blocks*2):
|
||||
block_sparse_offset = (tl.load(q_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_m
|
||||
qT = tl.load(qT_ptrs + block_sparse_offset * stride_tok)
|
||||
# Load m before computing qk to reduce pipeline stall.
|
||||
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
|
||||
m = tl.load(M + offs_m)
|
||||
qkT = tl.dot(k, qT)
|
||||
pT = tl.math.exp2(qkT - m[None, :])
|
||||
mask = tl.arange(0, BLOCK_N1) < block_size
|
||||
pT = tl.where(mask[:, None], pT, 0.0)
|
||||
|
||||
do = tl.load(do_ptrs + block_sparse_offset * stride_tok)
|
||||
# Compute dV.
|
||||
ppT = pT
|
||||
ppT = ppT.to(tl.bfloat16)
|
||||
dv += tl.dot(ppT, do)
|
||||
# D (= delta) is pre-divided by ds_scale.
|
||||
Di = tl.load(D + offs_m)
|
||||
# Compute dP and dS.
|
||||
dpT = tl.dot(v, tl.trans(do)).to(tl.float32)
|
||||
dsT = pT * (dpT - Di[None, :])
|
||||
dsT = dsT.to(tl.bfloat16)
|
||||
dk += tl.dot(dsT, tl.trans(qT))
|
||||
# Increment pointers.
|
||||
return dk, dv
|
||||
|
||||
|
||||
|
||||
# the main inner-loop logic for computing dQ
|
||||
@triton.jit
|
||||
def _attn_bwd_dq(dq, q, K, V, #
|
||||
do, m, D,
|
||||
# shared by Q/K/V/DO.
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M2: tl.constexpr, #
|
||||
BLOCK_N2: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr,
|
||||
# Filled in by the wrapper.
|
||||
start_m, start_n, num_steps):
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M2)
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N2)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
kT_ptrs = K + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
vT_ptrs = V + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
# D (= delta) is pre-divided by ds_scale.
|
||||
Di = tl.load(D + offs_m)
|
||||
# BLOCK_M2 must be a multiple of BLOCK_N2, otherwise the code wouldn't work.
|
||||
tl.static_assert(BLOCK_M2 % BLOCK_N2 == 0)
|
||||
step_n = BLOCK_N2
|
||||
|
||||
q_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(2) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_M2
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
|
||||
|
||||
for blk_idx in range(kv_blocks*2):
|
||||
block_sparse_offset = (tl.load(kv_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_n * stride_tok
|
||||
block_size = tl.load(variable_block_sizes + blk_idx//2) - (blk_idx%2) * step_n
|
||||
kT = tl.load(kT_ptrs + block_sparse_offset)
|
||||
vT = tl.load(vT_ptrs + block_sparse_offset)
|
||||
qk = tl.dot(q, kT)
|
||||
p = tl.math.exp2(qk - m)
|
||||
mask = tl.arange(0, BLOCK_N2) < block_size.to(tl.int32)
|
||||
p = tl.where(mask[None, :], p , 0.0)
|
||||
# Compute dP and dS.
|
||||
dp = tl.dot(do, vT).to(tl.float32)
|
||||
ds = p * (dp - Di[:, None])
|
||||
ds = ds.to(tl.bfloat16)
|
||||
# Compute dQ.
|
||||
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
|
||||
dq += tl.dot(ds, tl.trans(kT))
|
||||
# Increment pointers.
|
||||
return dq
|
||||
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd(Q, K, V, sm_scale, #
|
||||
DO, #
|
||||
DQ, DK, DV, #
|
||||
M, D,
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
# shared by Q/K/V/DO.
|
||||
stride_z, stride_h, stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M1: tl.constexpr, #
|
||||
BLOCK_N1: tl.constexpr, #
|
||||
BLOCK_M2: tl.constexpr, #
|
||||
BLOCK_N2: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr):
|
||||
LN2 = 0.6931471824645996 # = ln(2)
|
||||
|
||||
bhid = tl.program_id(2)
|
||||
off_chz = (bhid * N_CTX).to(tl.int64)
|
||||
adj = (stride_h * (bhid % H) + stride_z * (bhid // H)).to(tl.int64)
|
||||
pid = tl.program_id(0)
|
||||
|
||||
# offset pointers for batch/head
|
||||
Q += adj
|
||||
K += adj
|
||||
V += adj
|
||||
DO += adj
|
||||
DQ += adj
|
||||
DK += adj
|
||||
DV += adj
|
||||
M += off_chz
|
||||
D += off_chz
|
||||
|
||||
# load scales
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
|
||||
start_n = pid * BLOCK_N1
|
||||
start_m = 0
|
||||
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N1)
|
||||
|
||||
dv = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
|
||||
dk = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
|
||||
|
||||
# load K and V: they stay in SRAM throughout the inner loop.
|
||||
k = tl.load(K + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
v = tl.load(V + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
|
||||
|
||||
num_steps = N_CTX // BLOCK_M1
|
||||
|
||||
dk, dv = _attn_bwd_dkdv( #
|
||||
dk, dv, #
|
||||
Q, k, v, sm_scale, #
|
||||
DO, #
|
||||
M, D, #
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M1, BLOCK_N1, HEAD_DIM, #
|
||||
start_n, start_m, num_steps #
|
||||
)
|
||||
|
||||
dv_ptrs = DV + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
tl.store(dv_ptrs, dv)
|
||||
|
||||
# Write back dK.
|
||||
dk *= sm_scale
|
||||
dk_ptrs = DK + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
tl.store(dk_ptrs, dk)
|
||||
|
||||
# THIS BLOCK DOES DQ:
|
||||
start_m = pid * BLOCK_M2
|
||||
end_n = 0
|
||||
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M2)
|
||||
|
||||
q = tl.load(Q + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
dq = tl.zeros([BLOCK_M2, HEAD_DIM], dtype=tl.float32)
|
||||
do = tl.load(DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
|
||||
m = tl.load(M + offs_m)
|
||||
m = m[:, None]
|
||||
|
||||
num_steps = N_CTX // BLOCK_N2
|
||||
dq = _attn_bwd_dq(dq, q, K, V, #
|
||||
do, m, D, #
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M2, BLOCK_N2, HEAD_DIM, #
|
||||
start_m, end_n, num_steps #
|
||||
)
|
||||
# Write back dQ.
|
||||
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
dq *= LN2
|
||||
tl.store(dq_ptrs, dq)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num, variable_block_sizes):
|
||||
B, H, T, D = q.shape
|
||||
sm_scale = 1.0 / math.sqrt(D)
|
||||
max_kv_blks = q2k_index.shape[-1]
|
||||
assert T % 64 == 0, f"T must be a multiple of 64, but got {T}"
|
||||
assert T // 64 == q2k_num.shape[-1], f"shape mismatch, T // 64 = {T // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((B, H, T), dtype=torch.float32, device=q.device)
|
||||
|
||||
grid = lambda _: (triton.cdiv(T, 64), B * H, 1)
|
||||
_attn_fwd_sparse[grid](
|
||||
q, k, v, sm_scale,
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
M, o,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
|
||||
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
|
||||
v.stride(0), v.stride(1), v.stride(2), v.stride(3),
|
||||
o.stride(0), o.stride(1), o.stride(2), o.stride(3),
|
||||
B, H, T,
|
||||
HEAD_DIM=D, STAGE=3
|
||||
)
|
||||
|
||||
return o, M
|
||||
|
||||
def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q_index, k2q_num, variable_block_sizes):
|
||||
assert do.is_contiguous()
|
||||
assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
|
||||
|
||||
B, H, T, D = q.shape
|
||||
sm_scale = 1.0 / math.sqrt(D)
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
BATCH, N_HEAD, N_CTX = q.shape[:3]
|
||||
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
|
||||
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
|
||||
arg_k = k
|
||||
arg_k = arg_k * (sm_scale * RCP_LN2)
|
||||
PRE_BLOCK = 64
|
||||
assert N_CTX % PRE_BLOCK == 0
|
||||
pre_grid = (N_CTX // PRE_BLOCK, BATCH * N_HEAD)
|
||||
delta = torch.empty_like(M)
|
||||
_attn_bwd_preprocess[pre_grid](
|
||||
o, do, #
|
||||
delta, #
|
||||
BATCH, N_HEAD, N_CTX, #
|
||||
BLOCK_M=PRE_BLOCK, HEAD_DIM=D #
|
||||
)
|
||||
|
||||
|
||||
max_q_blks = k2q_index.shape[-1]
|
||||
max_kv_blks = q2k_index.shape[-1]
|
||||
|
||||
grid = (N_CTX // BLOCK_N1, 1, BATCH * N_HEAD)
|
||||
_attn_bwd[grid](
|
||||
q, arg_k, v, sm_scale, do, dq, dk, dv, #
|
||||
M, delta, #
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3), #
|
||||
N_HEAD, N_CTX, #
|
||||
BLOCK_M1=BLOCK_M1, BLOCK_N1=BLOCK_N1, #
|
||||
BLOCK_M2=BLOCK_M2, BLOCK_N2=BLOCK_N2, #
|
||||
HEAD_DIM=D #
|
||||
)
|
||||
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
@@ -1,185 +0,0 @@
|
||||
import torch
|
||||
try:
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
except ImportError:
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
from vsa.block_sparse_attn_triton import triton_block_sparse_attn_forward, triton_block_sparse_attn_backward
|
||||
assert torch.__version__ >= "2.4.0", "VSA requires PyTorch 2.4.0 or higher"
|
||||
from vsa.index import map_to_index
|
||||
from typing import Tuple, Optional
|
||||
|
||||
|
||||
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_triton", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_triton(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
block_map = block_map.int()
|
||||
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
|
||||
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
|
||||
return o, M
|
||||
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_triton")
|
||||
def _block_sparse_attn_triton_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
return o, M
|
||||
|
||||
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_backward_triton", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_backward_triton(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
M: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
|
||||
dq, dk, dv = triton_block_sparse_attn_backward(grad_output_padded, q_padded, k_padded, v_padded, o_padded, M, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes)
|
||||
return dq, dk, dv
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_backward_triton")
|
||||
def _block_sparse_attn_backward_triton_fake(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
M: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
dq = torch.empty_like(grad_output_padded)
|
||||
dk = torch.empty_like(grad_output_padded)
|
||||
dv = torch.empty_like(grad_output_padded)
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def backward_triton(ctx, grad_output1, grad_output2):
|
||||
q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_triton(grad_output1, q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
def setup_context_triton(ctx, inputs, output):
|
||||
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
|
||||
o_padded, M = output
|
||||
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes)
|
||||
|
||||
block_sparse_attn_triton.register_autograd(backward_triton, setup_context=setup_context_triton)
|
||||
|
||||
|
||||
major, minor = torch.cuda.get_device_capability(0)
|
||||
|
||||
if major == 9 and minor == 0:# check if H100
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_SM90", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_SM90(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
)-> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q_padded = q_padded.contiguous()
|
||||
k_padded = k_padded.contiguous()
|
||||
v_padded = v_padded.contiguous()
|
||||
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
|
||||
variable_block_sizes = variable_block_sizes.int()
|
||||
o_padded, lse_padded = block_sparse_fwd(q_padded, k_padded, v_padded, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
|
||||
return o_padded, lse_padded
|
||||
|
||||
|
||||
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_SM90")
|
||||
def _block_sparse_attn_SM90_fake(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q_padded, k_padded, v_padded = [x.contiguous() for x in (q_padded, k_padded, v_padded)]
|
||||
B, H, S, D = q_padded.shape
|
||||
o_padded = torch.empty_like(q_padded)
|
||||
lse_padded = torch.empty((B, H, S, 1), device=q_padded.device, dtype=torch.float32)
|
||||
return o_padded, lse_padded
|
||||
|
||||
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_backward_SM90", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_backward_SM90(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
)-> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
|
||||
grad_q_padded, grad_k_padded, grad_v_padded = block_sparse_bwd(
|
||||
q_padded, k_padded, v_padded, o_padded, lse_padded, grad_output_padded, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes
|
||||
)
|
||||
grad_q_padded = grad_q_padded.to(grad_output_padded.dtype)
|
||||
grad_k_padded = grad_k_padded.to(grad_output_padded.dtype)
|
||||
grad_v_padded = grad_v_padded.to(grad_output_padded.dtype)
|
||||
return grad_q_padded, grad_k_padded, grad_v_padded
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_backward_SM90")
|
||||
def _block_sparse_attn_backward_SM90_fake(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
torch._check(grad_output_padded.dtype == torch.bfloat16)
|
||||
torch._check(lse_padded.dtype == torch.float32)
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
dq = torch.empty_like(grad_output_padded)
|
||||
dk = torch.empty_like(grad_output_padded)
|
||||
dv = torch.empty_like(grad_output_padded)
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def backward_SM90(ctx, grad_output1, grad_output2):
|
||||
q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes= ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_SM90(grad_output1, q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
def setup_context_SM90(ctx, inputs, output):
|
||||
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
|
||||
o_padded, lse_padded = output
|
||||
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
|
||||
|
||||
|
||||
block_sparse_attn_SM90.register_autograd(backward_SM90, setup_context=setup_context_SM90)
|
||||
@@ -1,152 +0,0 @@
|
||||
|
||||
## pytorch sdpa version of block sparse ##
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
|
||||
@triton.jit
|
||||
def topk_index_to_map_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
topk,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
for i in tl.static_range(topk):
|
||||
index = tl.load(index_ptr_base + i * index_kv_stride)
|
||||
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
|
||||
|
||||
@triton.jit
|
||||
def map_to_index_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
index_num_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
index_num_bs_stride,
|
||||
index_num_h_stride,
|
||||
index_num_q_stride,
|
||||
num_kv_blocks,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
num = 0
|
||||
for i in tl.range(num_kv_blocks):
|
||||
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
|
||||
if map_entry:
|
||||
tl.store(index_ptr_base + num * index_kv_stride, i)
|
||||
num += 1
|
||||
|
||||
tl.store(
|
||||
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
|
||||
q * index_num_q_stride, num)
|
||||
|
||||
def topk_index_to_map(index: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
transpose_map: bool = False):
|
||||
"""
|
||||
Convert topk indices to a map.
|
||||
|
||||
Args:
|
||||
index: [bs, h, num_q_blocks, topk]
|
||||
The topk indices tensor.
|
||||
num_kv_blocks: int
|
||||
The number of key-value blocks in the block_map returned
|
||||
transpose_map: bool
|
||||
If True, the block_map will be transposed on the final two dimensions.
|
||||
|
||||
Returns:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
A binary map where 1 indicates that the q block attends to the kv block.
|
||||
"""
|
||||
bs, h, num_q_blocks, topk = index.shape
|
||||
|
||||
if transpose_map is False:
|
||||
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
else:
|
||||
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
block_map = block_map.transpose(2, 3)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
topk_index_to_map_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
topk=topk,
|
||||
)
|
||||
|
||||
return block_map
|
||||
|
||||
def map_to_index(block_map: torch.Tensor):
|
||||
"""
|
||||
Convert a block map to indices and counts.
|
||||
|
||||
Args:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The block map tensor.
|
||||
|
||||
Returns:
|
||||
index: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The indices of the blocks.
|
||||
index_num: [bs, h, num_q_blocks]
|
||||
The number of blocks for each q block.
|
||||
"""
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
|
||||
|
||||
index = torch.full((block_map.shape),
|
||||
-1,
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
index_num = torch.empty((bs, h, num_q_blocks),
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
map_to_index_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
index_num,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
index_num.stride(0),
|
||||
index_num.stride(1),
|
||||
index_num.stride(2),
|
||||
num_kv_blocks=num_kv_blocks,
|
||||
)
|
||||
|
||||
return index, index_num
|
||||
@@ -1,32 +0,0 @@
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## VMoBA: Mixture-of-Block Attention for Video Diffusion Models (VMoBA)
|
||||
|
||||
### Installation
|
||||
Please ensure that you have installed FlashAttention version **2.7.1 or higher**, as some interfaces have changed in recent releases.
|
||||
|
||||
### Usage
|
||||
|
||||
You can use `moba_attn_varlen` in the following ways:
|
||||
|
||||
**Install from source:**
|
||||
```bash
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
**Import after installation:**
|
||||
```python
|
||||
from vmoba import moba_attn_varlen
|
||||
```
|
||||
|
||||
**Or import directly from the project root:**
|
||||
```python
|
||||
from csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
|
||||
```
|
||||
|
||||
### Verify if you have successfully installed
|
||||
|
||||
```bash
|
||||
python csrc/attn/vmoba_attn/vmoba/vmoba.py
|
||||
```
|
||||
|
||||
@@ -1,24 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from setuptools import find_packages, setup
|
||||
|
||||
PACKAGE_NAME = "vmoba"
|
||||
VERSION = "0.0.0"
|
||||
AUTHOR = "JianzongWu"
|
||||
DESCRIPTION = "VMoBA: Mixture-of-Block Attention for Video Diffusion Models"
|
||||
URL = "https://github.com/KwaiVGI/VMoBA"
|
||||
|
||||
setup(
|
||||
name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
classifiers=[
|
||||
"Programming Language :: Python :: 3",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
],
|
||||
python_requires='>=3.12',
|
||||
install_requires=[]
|
||||
)
|
||||
@@ -1,97 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
import pytest
|
||||
import random
|
||||
from csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
|
||||
|
||||
def generate_test_data(batch_size, total_seqlen, num_heads, head_dim, dtype, device="cuda"):
|
||||
"""
|
||||
Generates random data for testing the variable-length attention function.
|
||||
"""
|
||||
torch.manual_seed(42)
|
||||
random.seed(42)
|
||||
torch.cuda.manual_seed_all(42)
|
||||
|
||||
# Generate sequence lengths for each item in the batch
|
||||
if batch_size > 1:
|
||||
# Ensure sequence lengths are reasonably distributed
|
||||
avg_seqlen = total_seqlen // batch_size
|
||||
seqlens = [random.randint(avg_seqlen // 2, avg_seqlen + avg_seqlen // 2) for _ in range(batch_size - 1)]
|
||||
remaining_len = total_seqlen - sum(seqlens)
|
||||
if remaining_len > 0:
|
||||
seqlens.append(remaining_len)
|
||||
else: # Adjust if sum exceeds total_seqlen
|
||||
seqlens.append(avg_seqlen)
|
||||
current_sum = sum(seqlens)
|
||||
seqlens[-1] -= (current_sum - total_seqlen)
|
||||
# Ensure all lengths are positive
|
||||
seqlens = [max(1, s) for s in seqlens]
|
||||
# Final adjustment to match total_seqlen
|
||||
seqlens[-1] += total_seqlen - sum(seqlens)
|
||||
|
||||
else:
|
||||
seqlens = [total_seqlen]
|
||||
|
||||
cu_seqlens = torch.tensor([0] + list(torch.cumsum(torch.tensor(seqlens), 0)), device=device, dtype=torch.int32)
|
||||
max_seqlen = max(seqlens) if seqlens else 0
|
||||
|
||||
q = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
k = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
v = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
|
||||
|
||||
return q, k, v, cu_seqlens, max_seqlen
|
||||
|
||||
|
||||
@pytest.mark.parametrize("batch_size", [1, 2])
|
||||
@pytest.mark.parametrize("total_seqlen", [512, 1024])
|
||||
@pytest.mark.parametrize("num_heads", [8])
|
||||
@pytest.mark.parametrize("head_dim", [64])
|
||||
@pytest.mark.parametrize("moba_chunk_size", [64])
|
||||
@pytest.mark.parametrize("moba_topk", [2, 4])
|
||||
@pytest.mark.parametrize("select_mode", ["topk", "threshold"])
|
||||
@pytest.mark.parametrize("threshold_type", ["query_head", "head_global", "overall"])
|
||||
@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16])
|
||||
def test_moba_attn_varlen_forward(
|
||||
batch_size, total_seqlen, num_heads, head_dim, moba_chunk_size, moba_topk, select_mode, threshold_type, dtype
|
||||
):
|
||||
"""
|
||||
Tests the forward pass of moba_attn_varlen for basic correctness.
|
||||
It checks output shape, dtype, and for the presence of NaNs/Infs.
|
||||
"""
|
||||
if dtype == torch.float32:
|
||||
pytest.skip("float32 is not supported in flash attention")
|
||||
|
||||
q, k, v, cu_seqlens, max_seqlen = generate_test_data(
|
||||
batch_size, total_seqlen, num_heads, head_dim, dtype
|
||||
)
|
||||
|
||||
# Ensure chunk size is not larger than the smallest sequence length
|
||||
min_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).min().item()
|
||||
if moba_chunk_size > min_seqlen:
|
||||
pytest.skip("moba_chunk_size is larger than the minimum sequence length in the batch")
|
||||
|
||||
try:
|
||||
output = moba_attn_varlen(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens=cu_seqlens,
|
||||
max_seqlen=max_seqlen,
|
||||
moba_chunk_size=moba_chunk_size,
|
||||
moba_topk=moba_topk,
|
||||
select_mode=select_mode,
|
||||
threshold_type=threshold_type,
|
||||
simsum_threshold=0.5, # A reasonable default for threshold mode
|
||||
)
|
||||
except Exception as e:
|
||||
pytest.fail(f"moba_attn_varlen forward pass failed with exception: {e}")
|
||||
|
||||
# 1. Check output shape
|
||||
assert output.shape == q.shape, f"Expected output shape {q.shape}, but got {output.shape}"
|
||||
|
||||
# 2. Check output dtype
|
||||
assert output.dtype == q.dtype, f"Expected output dtype {q.dtype}, but got {output.dtype}"
|
||||
|
||||
# 3. Check for NaNs or Infs in the output
|
||||
assert torch.all(torch.isfinite(output)), "Output contains NaN or Inf values"
|
||||
@@ -1,2 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from .vmoba import moba_attn_varlen, process_moba_input, process_moba_output
|
||||
@@ -1,860 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapt from https://github.com/KwaiVGI/VMoBA/blob/main/src/vmoba.py
|
||||
|
||||
import random
|
||||
import time
|
||||
import os
|
||||
import torch
|
||||
from typing import Tuple
|
||||
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
|
||||
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
|
||||
from functools import lru_cache
|
||||
from einops import rearrange
|
||||
|
||||
|
||||
@lru_cache(maxsize=16)
|
||||
def calc_chunks(cu_seqlen, moba_chunk_size):
|
||||
"""
|
||||
Calculate chunk boundaries.
|
||||
|
||||
For vision tasks we include all chunks (even the last one which might be shorter)
|
||||
so that every chunk can be selected.
|
||||
"""
|
||||
batch_sizes = cu_seqlen[1:] - cu_seqlen[:-1]
|
||||
batch_num_chunk = (batch_sizes + (moba_chunk_size - 1)) // moba_chunk_size
|
||||
cu_num_chunk = torch.ones(
|
||||
batch_num_chunk.numel() + 1,
|
||||
device=cu_seqlen.device,
|
||||
dtype=batch_num_chunk.dtype,
|
||||
)
|
||||
cu_num_chunk[1:] = batch_num_chunk.cumsum(dim=0)
|
||||
num_chunk = cu_num_chunk[-1]
|
||||
chunk_sizes = torch.full(
|
||||
(num_chunk + 1,), moba_chunk_size, dtype=torch.int32, device=cu_seqlen.device
|
||||
)
|
||||
chunk_sizes[0] = 0
|
||||
batch_last_chunk_size = batch_sizes - (batch_num_chunk - 1) * moba_chunk_size
|
||||
chunk_sizes[cu_num_chunk[1:]] = batch_last_chunk_size
|
||||
cu_chunk = chunk_sizes.cumsum(dim=-1, dtype=torch.int32)
|
||||
chunk_to_batch = torch.zeros(
|
||||
(num_chunk,), dtype=torch.int32, device=cu_seqlen.device
|
||||
)
|
||||
chunk_to_batch[cu_num_chunk[1:-1]] = 1
|
||||
chunk_to_batch = chunk_to_batch.cumsum(dim=0, dtype=torch.int32)
|
||||
|
||||
# Do not filter out any chunk
|
||||
filtered_chunk_indices = torch.arange(
|
||||
num_chunk, device=cu_seqlen.device, dtype=torch.int32
|
||||
)
|
||||
num_filtered_chunk = num_chunk
|
||||
|
||||
return cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch
|
||||
|
||||
|
||||
# --- Threshold Selection Helper Functions ---
|
||||
|
||||
def _select_threshold_query_head(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects chunks for each <query, head> pair based on threshold.
|
||||
Normalization and sorting happen along the chunk dimension (dim=0).
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization per <head, query> (across chunks)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
row_min = gate_min_val.amin(dim=0) # (H, S)
|
||||
row_max = gate_masked.amax(dim=0) # (H, S)
|
||||
denom = row_max - row_min
|
||||
denom = torch.where(denom <= eps, torch.ones_like(denom), denom) # avoid divide‑by‑zero
|
||||
|
||||
gate_norm = (gate - row_min.unsqueeze(0)) / denom.unsqueeze(0)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) pull out the self‐chunk’s normalized weight for each <head,seq>
|
||||
self_norm = (gate_norm * gate_self_chunk_mask).sum(dim=0) # (H, S)
|
||||
|
||||
# 2) compute how much more normalized weight we need beyond self
|
||||
total_norm_sum = gate_norm.sum(dim=0) # (H, S)
|
||||
remain_ratio = simsum_threshold - self_norm / (total_norm_sum + eps) # (H, S)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # if already ≥ thresh, no extra needed
|
||||
|
||||
# 3) zero out the self‐chunk in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0
|
||||
|
||||
# 4) sort the other chunks by descending norm, per <head,seq>
|
||||
sorted_norm, sorted_idx = torch.sort(others_norm, descending=True, dim=0) # (C, H, S)
|
||||
|
||||
# 5) cumulative‑sum the sorted norms per <head,seq>
|
||||
cumsum_others = sorted_norm.cumsum(dim=0) # (C, H, S)
|
||||
|
||||
# 6) for each <head,seq>, find the smallest k where cumsum_ratio ≥ remain_ratio
|
||||
ratio = cumsum_others / (total_norm_sum.unsqueeze(0) + eps) # (C, H, S)
|
||||
cond = ratio >= remain_ratio.unsqueeze(0) # (C, H, S) boolean mask
|
||||
any_cond = cond.any(dim=0) # (H, S)
|
||||
# Find the index of the first True value along dim 0. If none, use C-1.
|
||||
cutoff = torch.where(any_cond, cond.float().argmax(dim=0), torch.full_like(any_cond, fill_value=C - 1)) # (H, S)
|
||||
|
||||
# 7) build a mask in sorted order up to that cutoff
|
||||
idx_range = torch.arange(C, device=gate.device).view(-1, 1, 1) # (C, 1, 1)
|
||||
sorted_mask = idx_range <= cutoff.unsqueeze(0) # (C, H, S)
|
||||
|
||||
# 8) scatter it back to original chunk order
|
||||
others_mask = torch.zeros_like(gate, dtype=torch.bool)
|
||||
others_mask.scatter_(0, sorted_idx, sorted_mask)
|
||||
|
||||
# 9) finally, include every self‐chunk plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_block(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <query, head> pairs for each block based on threshold.
|
||||
Normalization and sorting happen across the head and sequence dimensions (dim=1, 2).
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
HS = H * S
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization per block (across heads and queries)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
block_max = gate_masked.amax(dim=(1, 2), keepdim=True) # (C, 1, 1)
|
||||
block_min = gate_min_val.amin(dim=(1, 2), keepdim=True) # (C, 1, 1)
|
||||
block_denom = block_max - block_min
|
||||
block_denom = torch.where(block_denom <= eps, torch.ones_like(block_denom), block_denom) # (C, 1, 1)
|
||||
|
||||
gate_norm = (gate - block_min) / block_denom # (C, H, S)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) identify normalized weights of entries that *are* self-chunks (from query perspective)
|
||||
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
|
||||
# Sum these weights *per block*
|
||||
self_norm_sum_per_block = self_norm_entries.sum(dim=(1, 2)) # (C,)
|
||||
|
||||
# 2) compute how much more normalized weight each block needs beyond its self-chunk contributions
|
||||
total_norm_sum_per_block = gate_norm.sum(dim=(1, 2)) # (C,)
|
||||
remain_ratio = simsum_threshold - self_norm_sum_per_block / (total_norm_sum_per_block + eps) # (C,)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # (C,)
|
||||
|
||||
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
|
||||
|
||||
# 4) sort the other <head, seq> pairs by descending norm, per block
|
||||
others_flat = others_norm.contiguous().view(C, HS) # (C, H*S)
|
||||
sorted_others_flat, sorted_indices_flat = torch.sort(others_flat, dim=1, descending=True) # (C, H*S)
|
||||
|
||||
# 5) cumulative‑sum the sorted norms per block
|
||||
cumsum_others_flat = sorted_others_flat.cumsum(dim=1) # (C, H*S)
|
||||
|
||||
# 6) for each block, find the smallest k where cumsum_ratio ≥ remain_ratio
|
||||
ratio_flat = cumsum_others_flat / (total_norm_sum_per_block.unsqueeze(1) + eps) # (C, H*S)
|
||||
cond_flat = ratio_flat >= remain_ratio.unsqueeze(1) # (C, H*S) boolean mask
|
||||
any_cond = cond_flat.any(dim=1) # (C,)
|
||||
# Find the index of the first True value along dim 1. If none, use HS-1.
|
||||
cutoff_flat = torch.where(any_cond, cond_flat.float().argmax(dim=1), torch.full_like(any_cond, fill_value=HS - 1)) # (C,)
|
||||
|
||||
# 7) build a mask in sorted order up to that cutoff per block
|
||||
idx_range_flat = torch.arange(HS, device=gate.device).unsqueeze(0) # (1, H*S)
|
||||
sorted_mask_flat = idx_range_flat <= cutoff_flat.unsqueeze(1) # (C, H*S)
|
||||
|
||||
# 8) scatter it back to original <head, seq> order per block
|
||||
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C, H*S)
|
||||
others_mask_flat.scatter_(1, sorted_indices_flat, sorted_mask_flat)
|
||||
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
|
||||
|
||||
# 9) finally, include every self‐chunk entry plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_overall(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <chunk, query, head> triplets globally based on threshold.
|
||||
Normalization and sorting happen across all valid entries.
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
CHS = C * H * S
|
||||
eps = 1e-6
|
||||
|
||||
# LSE‐style normalization globally across all valid entries
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
|
||||
|
||||
overall_max = gate_masked.max() # scalar
|
||||
overall_min = gate_min_val.min() # scalar
|
||||
overall_denom = overall_max - overall_min
|
||||
overall_denom = torch.where(overall_denom <= eps, torch.tensor(1.0, device=gate.device, dtype=gate.dtype), overall_denom)
|
||||
|
||||
gate_norm = (gate - overall_min) / overall_denom # (C, H, S)
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 1) identify normalized weights of entries that *are* self-chunks
|
||||
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
|
||||
# Sum these weights globally
|
||||
self_norm_sum_overall = self_norm_entries.sum() # scalar
|
||||
|
||||
# 2) compute how much more normalized weight is needed globally beyond self-chunk contributions
|
||||
total_norm_sum_overall = gate_norm.sum() # scalar
|
||||
remain_ratio = simsum_threshold - self_norm_sum_overall / (total_norm_sum_overall + eps) # scalar
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0) # scalar
|
||||
|
||||
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
|
||||
|
||||
# 4) sort all other entries by descending norm, globally
|
||||
others_flat = others_norm.flatten() # (C*H*S,)
|
||||
valid_others_mask_flat = valid_gate_mask.flatten() & ~gate_self_chunk_mask.flatten() # Mask for valid, non-self entries
|
||||
|
||||
# Only sort the valid 'other' entries
|
||||
valid_others_indices = torch.where(valid_others_mask_flat)[0]
|
||||
valid_others_values = others_flat[valid_others_indices]
|
||||
|
||||
sorted_others_values, sort_perm = torch.sort(valid_others_values, descending=True) # (N_valid_others,)
|
||||
sorted_original_indices = valid_others_indices[sort_perm] # Original indices in C*H*S space, sorted by value
|
||||
|
||||
# 5) cumulative‑sum the sorted valid 'other' norms globally
|
||||
cumsum_others_values = sorted_others_values.cumsum(dim=0) # (N_valid_others,)
|
||||
|
||||
# 6) find the smallest k where cumsum_ratio ≥ remain_ratio globally
|
||||
ratio_values = cumsum_others_values / (total_norm_sum_overall + eps) # (N_valid_others,)
|
||||
cond_values = ratio_values >= remain_ratio # (N_valid_others,) boolean mask
|
||||
any_cond = cond_values.any() # scalar
|
||||
|
||||
# Find the index of the first True value in the *sorted* list. If none, use all valid others.
|
||||
cutoff_idx_in_sorted = torch.where(
|
||||
any_cond,
|
||||
cond_values.float().argmax(dim=0),
|
||||
torch.tensor(len(sorted_others_values) - 1, device=gate.device, dtype=torch.long)
|
||||
)
|
||||
|
||||
# 7) build a mask selecting the top-k others based on the cutoff
|
||||
# Select the original indices corresponding to the top entries in the sorted list
|
||||
selected_other_indices = sorted_original_indices[:cutoff_idx_in_sorted + 1]
|
||||
|
||||
# 8) create the mask in the original flat shape
|
||||
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C*H*S,)
|
||||
if selected_other_indices.numel() > 0: # Check if any 'other' indices were selected
|
||||
others_mask_flat[selected_other_indices] = True
|
||||
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
|
||||
|
||||
# 9) finally, include every self‐chunk entry plus all selected others
|
||||
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
def _select_threshold_head_global(
|
||||
gate: torch.Tensor,
|
||||
valid_gate_mask: torch.Tensor,
|
||||
gate_self_chunk_mask: torch.Tensor,
|
||||
simsum_threshold: float
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Selects <chunk, query> globally for each head based on threshold.
|
||||
"""
|
||||
C, H, S = gate.shape
|
||||
eps = 1e-6
|
||||
|
||||
# 1) LSE‐style normalization per head (across chunks and sequence dims)
|
||||
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf)
|
||||
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf)
|
||||
|
||||
max_per_head = gate_masked.amax(dim=(0, 2), keepdim=True) # (1, H, 1)
|
||||
min_per_head = gate_min_val.amin(dim=(0, 2), keepdim=True) # (1, H, 1)
|
||||
denom = max_per_head - min_per_head
|
||||
denom = torch.where(denom <= eps, torch.ones_like(denom), denom)
|
||||
|
||||
gate_norm = (gate - min_per_head) / denom
|
||||
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
|
||||
|
||||
# 2) sum normalized self‐chunk contributions per head
|
||||
self_norm_sum = (gate_norm * gate_self_chunk_mask).sum(dim=(0, 2)) # (H,)
|
||||
|
||||
# 3) total normalized sum per head
|
||||
total_norm_sum = gate_norm.sum(dim=(0, 2)) # (H,)
|
||||
|
||||
# 4) how much more normalized weight needed per head
|
||||
remain_ratio = simsum_threshold - self_norm_sum / (total_norm_sum + eps) # (H,)
|
||||
remain_ratio = torch.clamp(remain_ratio, min=0.0)
|
||||
|
||||
# 5) zero out self‐chunk entries to focus on "others"
|
||||
others_norm = gate_norm.clone()
|
||||
others_norm[gate_self_chunk_mask] = 0.0 # (C, H, S)
|
||||
|
||||
# 6) flatten chunk and sequence dims, per head
|
||||
CS = C * S
|
||||
others_flat = others_norm.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
|
||||
valid_flat = (valid_gate_mask & ~gate_self_chunk_mask) \
|
||||
.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
|
||||
|
||||
# 7) vectorized selection of “others” per head
|
||||
masked_flat = torch.where(valid_flat, others_flat, torch.zeros_like(others_flat))
|
||||
sorted_vals, sorted_idx = torch.sort(masked_flat, dim=1, descending=True) # (H, C*S)
|
||||
|
||||
cumsum_vals = sorted_vals.cumsum(dim=1) # (H, C*S)
|
||||
ratio_vals = cumsum_vals / (total_norm_sum.unsqueeze(1) + eps) # (H, C*S)
|
||||
cond = ratio_vals >= remain_ratio.unsqueeze(1) # (H, C*S)
|
||||
|
||||
has_cutoff = cond.any(dim=1) # (H,)
|
||||
default = torch.full((H,), CS - 1, device=gate.device, dtype=torch.long)
|
||||
cutoff = torch.where(has_cutoff, cond.float().argmax(dim=1), default) # (H,)
|
||||
|
||||
idx_range = torch.arange(CS, device=gate.device).unsqueeze(0) # (1, C*S)
|
||||
sorted_mask = idx_range <= cutoff.unsqueeze(1) # (H, C*S)
|
||||
|
||||
selected_flat = torch.zeros_like(valid_flat) # (H, C*S)
|
||||
selected_flat.scatter_(1, sorted_idx, sorted_mask) # (H, C*S)
|
||||
|
||||
# 8) reshape selection mask back to (C, H, S)
|
||||
others_mask = selected_flat.reshape(H, C, S).permute(1, 0, 2) # (C, H, S)
|
||||
|
||||
# 9) include self‐chunks plus selected others, and obey valid mask
|
||||
final_gate_mask = valid_gate_mask & (gate_self_chunk_mask | others_mask)
|
||||
|
||||
return final_gate_mask
|
||||
|
||||
|
||||
class MixedAttention(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(
|
||||
ctx,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
max_seqlen,
|
||||
moba_chunk_size,
|
||||
moba_q_sh_indices,
|
||||
):
|
||||
ctx.max_seqlen = max_seqlen
|
||||
ctx.moba_chunk_size = moba_chunk_size
|
||||
ctx.softmax_scale = softmax_scale = q.shape[-1] ** (-0.5)
|
||||
|
||||
# Non-causal self-attention branch
|
||||
# return out, softmax_lse, S_dmask, rng_state
|
||||
self_attn_out_sh, self_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
cu_seqlens_q=self_attn_cu_seqlen,
|
||||
cu_seqlens_k=self_attn_cu_seqlen,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=max_seqlen,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
)
|
||||
# MOBA attention branch (non-causal)
|
||||
moba_attn_out, moba_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
|
||||
q=moba_q,
|
||||
k=moba_kv[:, 0],
|
||||
v=moba_kv[:, 1],
|
||||
cu_seqlens_q=moba_cu_seqlen_q,
|
||||
cu_seqlens_k=moba_cu_seqlen_kv,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=moba_chunk_size,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
)
|
||||
|
||||
self_attn_lse_sh = self_attn_lse_hs.t().contiguous()
|
||||
moba_attn_lse = moba_attn_lse_hs.t().contiguous()
|
||||
|
||||
output = torch.zeros((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
output_2d = output.view(-1, q.shape[2])
|
||||
|
||||
max_lse_1d = self_attn_lse_sh.view(-1)
|
||||
max_lse_1d = max_lse_1d.index_reduce(
|
||||
0, moba_q_sh_indices, moba_attn_lse.view(-1), "amax"
|
||||
)
|
||||
self_attn_lse_sh = self_attn_lse_sh - max_lse_1d.view_as(self_attn_lse_sh)
|
||||
moba_attn_lse = (
|
||||
moba_attn_lse.view(-1)
|
||||
.sub(max_lse_1d.index_select(0, moba_q_sh_indices))
|
||||
.reshape_as(moba_attn_lse)
|
||||
)
|
||||
|
||||
mixed_attn_se_sh = self_attn_lse_sh.exp()
|
||||
moba_attn_se = moba_attn_lse.exp()
|
||||
|
||||
mixed_attn_se_sh.view(-1).index_add_(
|
||||
0, moba_q_sh_indices, moba_attn_se.view(-1)
|
||||
)
|
||||
mixed_attn_lse_sh = mixed_attn_se_sh.log()
|
||||
|
||||
# Combine self-attention output
|
||||
factor = (self_attn_lse_sh - mixed_attn_lse_sh).exp() # [S, H]
|
||||
self_attn_out_sh = self_attn_out_sh * factor.unsqueeze(-1)
|
||||
output_2d += self_attn_out_sh.reshape_as(output_2d)
|
||||
|
||||
# Combine MOBA attention output
|
||||
mixed_attn_lse = (
|
||||
mixed_attn_lse_sh.view(-1)
|
||||
.index_select(0, moba_q_sh_indices)
|
||||
.view_as(moba_attn_lse)
|
||||
)
|
||||
factor = (moba_attn_lse - mixed_attn_lse).exp() # [S, H]
|
||||
moba_attn_out = moba_attn_out * factor.unsqueeze(-1)
|
||||
raw_attn_out = moba_attn_out.view(-1, moba_attn_out.shape[-1])
|
||||
output_2d.index_add_(0, moba_q_sh_indices, raw_attn_out)
|
||||
output = output.to(q.dtype)
|
||||
mixed_attn_lse_sh = mixed_attn_lse_sh + max_lse_1d.view_as(mixed_attn_se_sh)
|
||||
ctx.save_for_backward(
|
||||
output,
|
||||
mixed_attn_lse_sh,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
moba_q_sh_indices,
|
||||
)
|
||||
|
||||
return output
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, d_output):
|
||||
|
||||
max_seqlen = ctx.max_seqlen
|
||||
moba_chunk_size = ctx.moba_chunk_size
|
||||
softmax_scale = ctx.softmax_scale
|
||||
|
||||
(
|
||||
output,
|
||||
mixed_attn_vlse_sh,
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
moba_q_sh_indices,
|
||||
) = ctx.saved_tensors
|
||||
|
||||
d_output = d_output.contiguous()
|
||||
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
_ = _flash_attn_varlen_backward(
|
||||
dout=d_output,
|
||||
q=q,
|
||||
k=k,
|
||||
v=v,
|
||||
out=output,
|
||||
softmax_lse=mixed_attn_vlse_sh.t().contiguous(),
|
||||
dq=dq,
|
||||
dk=dk,
|
||||
dv=dv,
|
||||
cu_seqlens_q=self_attn_cu_seqlen,
|
||||
cu_seqlens_k=self_attn_cu_seqlen,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=max_seqlen,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=True,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1
|
||||
)
|
||||
|
||||
headdim = q.shape[-1]
|
||||
d_moba_output = (
|
||||
d_output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
|
||||
)
|
||||
moba_output = (
|
||||
output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
|
||||
)
|
||||
|
||||
mixed_attn_vlse = (
|
||||
mixed_attn_vlse_sh.view(-1).index_select(0, moba_q_sh_indices).view(1, -1)
|
||||
)
|
||||
|
||||
dmq = torch.empty_like(moba_q)
|
||||
dmkv = torch.empty_like(moba_kv)
|
||||
_ = _flash_attn_varlen_backward(
|
||||
dout=d_moba_output,
|
||||
q=moba_q,
|
||||
k=moba_kv[:, 0],
|
||||
v=moba_kv[:, 1],
|
||||
out=moba_output,
|
||||
softmax_lse=mixed_attn_vlse,
|
||||
dq=dmq,
|
||||
dk=dmkv[:,0],
|
||||
dv=dmkv[:,1],
|
||||
cu_seqlens_q=moba_cu_seqlen_q,
|
||||
cu_seqlens_k=moba_cu_seqlen_kv,
|
||||
max_seqlen_q=max_seqlen,
|
||||
max_seqlen_k=moba_chunk_size,
|
||||
softmax_scale=softmax_scale,
|
||||
causal=False,
|
||||
dropout_p=0.0,
|
||||
softcap=0.0,
|
||||
alibi_slopes=None,
|
||||
deterministic=True,
|
||||
window_size_left=-1,
|
||||
window_size_right=-1
|
||||
)
|
||||
|
||||
return dq, dk, dv, None, dmq, dmkv, None, None, None, None, None
|
||||
|
||||
|
||||
def moba_attn_varlen(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
cu_seqlens: torch.Tensor,
|
||||
max_seqlen: int,
|
||||
moba_chunk_size: int,
|
||||
moba_topk: int,
|
||||
select_mode: str = 'threshold', # "topk" or "threshold"
|
||||
simsum_threshold: float = 0.25,
|
||||
threshold_type: str = 'query_head',
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Accelerated MOBA attention for vision tasks with proper LSE normalization.
|
||||
|
||||
This version:
|
||||
- Splits KV into chunks.
|
||||
- For each query head, selects the top-k relevant KV chunks (including the self chunk)
|
||||
by amplifying the diagonal (self-chunk) logits.
|
||||
- Aggregates the attention outputs from the selected chunks using a log-sum-exp
|
||||
reduction so that attending to each query over the selected chunks is equivalent
|
||||
to the original algorithm.
|
||||
"""
|
||||
# Stack keys and values.
|
||||
kv = torch.stack((k, v), dim=1)
|
||||
seqlen, num_head, head_dim = q.shape
|
||||
|
||||
# Compute chunk boundaries.
|
||||
cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch = calc_chunks(
|
||||
cu_seqlens, moba_chunk_size
|
||||
)
|
||||
|
||||
self_attn_cu_seqlen = cu_chunk
|
||||
|
||||
# Update top-k selection to include the self chunk.
|
||||
moba_topk = min(moba_topk, num_filtered_chunk)
|
||||
|
||||
# --- Build filtered KV from chunks ---
|
||||
chunk_starts = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
|
||||
chunk_ends = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
|
||||
chunk_lengths = chunk_ends - chunk_starts # [num_filtered_chunk]
|
||||
max_chunk_len = int(chunk_lengths.max().item())
|
||||
|
||||
range_tensor = torch.arange(max_chunk_len, device=kv.device, dtype=chunk_starts.dtype).unsqueeze(0)
|
||||
indices = chunk_starts.unsqueeze(1) + range_tensor
|
||||
indices = torch.clamp(indices, max=kv.shape[0] - 1)
|
||||
valid_mask = range_tensor < chunk_lengths.unsqueeze(1)
|
||||
gathered = kv[indices.view(-1)].view(num_filtered_chunk, max_chunk_len, *kv.shape[1:])
|
||||
gathered = gathered * valid_mask.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1).type_as(gathered)
|
||||
|
||||
# Compute key_gate_weight over valid tokens.
|
||||
key_values = gathered[:, :, 0].float() # [num_filtered_chunk, max_chunk_len, num_head, head_dim]
|
||||
valid_mask_exp = valid_mask.unsqueeze(-1).unsqueeze(-1)
|
||||
key_sum = (key_values * valid_mask_exp).sum(dim=1)
|
||||
divisor = valid_mask.sum(dim=1).unsqueeze(-1).unsqueeze(-1)
|
||||
key_gate_weight = key_sum / divisor # [num_filtered_chunk, num_head, head_dim]
|
||||
|
||||
# Compute gate logits between key_gate_weight and queries.
|
||||
q_float = q.float()
|
||||
# gate = torch.einsum("nhd,shd->nhs", key_gate_weight, q_float) # [num_filtered_chunk, num_head, seqlen]
|
||||
gate = torch.bmm(key_gate_weight.permute(1, 0, 2), q_float.permute(1, 0, 2).transpose(1, 2)).permute(1, 0, 2)
|
||||
|
||||
# Amplify the diagonal (self chunk) contributions.
|
||||
gate_seq_idx = torch.arange(seqlen, device=q.device, dtype=torch.int32).unsqueeze(0).expand(num_filtered_chunk, seqlen)
|
||||
chunk_start = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
|
||||
chunk_end = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
|
||||
gate_self_chunk_mask = ((gate_seq_idx >= chunk_start.unsqueeze(1)) &
|
||||
(gate_seq_idx < chunk_end.unsqueeze(1))).unsqueeze(1).expand(-1, num_head, -1)
|
||||
amplification_factor = 1e9 # Example factor; adjust as needed.
|
||||
origin_gate = gate.clone()
|
||||
gate = gate.clone()
|
||||
if select_mode == "topk":
|
||||
gate[gate_self_chunk_mask] += amplification_factor
|
||||
|
||||
# Exclude positions that are outside the valid batch boundaries.
|
||||
batch_starts = cu_seqlens[chunk_to_batch[filtered_chunk_indices]]
|
||||
batch_ends = cu_seqlens[chunk_to_batch[filtered_chunk_indices] + 1]
|
||||
gate_batch_start_mask = gate_seq_idx < batch_starts.unsqueeze(1)
|
||||
gate_batch_end_mask = gate_seq_idx >= batch_ends.unsqueeze(1)
|
||||
gate_inf_mask = gate_batch_start_mask | gate_batch_end_mask
|
||||
gate.masked_fill_(gate_inf_mask.unsqueeze(1), -float("inf"))
|
||||
|
||||
if select_mode == 'topk':
|
||||
# We amplify self‐chunk in gate already, so self entries will rank highest.
|
||||
valid_gate_mask = gate != -float("inf")
|
||||
if threshold_type == 'query_head':
|
||||
# === per‐<head,seq> top-k across chunks (original behavior) ===
|
||||
# gate: (C, H, S)
|
||||
_, gate_topk_idx = torch.topk(gate, k=moba_topk, dim=0, largest=True, sorted=False)
|
||||
gate_idx_mask = torch.zeros_like(gate, dtype=torch.bool)
|
||||
gate_idx_mask.scatter_(0, gate_topk_idx, True)
|
||||
gate_mask = valid_gate_mask & gate_idx_mask
|
||||
elif threshold_type == 'overall':
|
||||
# === global top-k across all (chunk, head, seq) entries ===
|
||||
C, H, S = gate.shape
|
||||
flat_gate = gate.flatten()
|
||||
flat_mask = valid_gate_mask.flatten()
|
||||
flat_gate_masked = torch.where(flat_mask, flat_gate, -float("inf"))
|
||||
# pick topk global entries
|
||||
vals, idx = torch.topk(flat_gate_masked, k=moba_topk * H * S, largest=True, sorted=False)
|
||||
others_mask_flat = torch.zeros_like(flat_mask, dtype=torch.bool)
|
||||
others_mask_flat[idx] = True
|
||||
gate_mask = (valid_gate_mask.flatten() & others_mask_flat).view(gate.shape)
|
||||
elif threshold_type == 'head_global':
|
||||
# per-head top-k across all chunks and sequence positions
|
||||
C, H, S = gate.shape
|
||||
CS = C * S
|
||||
flat_gate = gate.permute(1, 0, 2).reshape(H, CS)
|
||||
flat_valid = valid_gate_mask.permute(1, 0, 2).reshape(H, CS)
|
||||
flat_gate_masked = torch.where(flat_valid, flat_gate, torch.full_like(flat_gate, -float('inf')))
|
||||
# pick top-k indices per head
|
||||
_, topk_idx = torch.topk(flat_gate_masked, k=moba_topk * S, dim=1, largest=True, sorted=False)
|
||||
gate_idx_flat = torch.zeros_like(flat_valid, dtype=torch.bool)
|
||||
gate_idx_flat.scatter_(1, topk_idx, True)
|
||||
gate_mask = gate_idx_flat.reshape(H, C, S).permute(1, 0, 2)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Invalid threshold_type for topk: {threshold_type}. "
|
||||
"Choose 'query_head', 'block', or 'overall'."
|
||||
)
|
||||
elif select_mode == 'threshold':
|
||||
# Delegate to the specific thresholding function
|
||||
valid_gate_mask = gate != -float("inf") # (num_chunk, num_head, seqlen)
|
||||
if threshold_type == 'query_head':
|
||||
gate_mask = _select_threshold_query_head(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'block':
|
||||
gate_mask = _select_threshold_block(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'overall':
|
||||
gate_mask = _select_threshold_overall(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
elif threshold_type == 'head_global':
|
||||
gate_mask = _select_threshold_head_global(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
|
||||
else:
|
||||
raise ValueError(f"Invalid threshold_type: {threshold_type}. Choose 'query_head', 'block', or 'overall'.")
|
||||
else:
|
||||
raise ValueError(f"Invalid select_mode: {select_mode}. Choose 'topk' or 'threshold'.")
|
||||
|
||||
# eliminate self_chunk in MoBA branch
|
||||
gate_mask = gate_mask & ~gate_self_chunk_mask
|
||||
# if gate_mask is all false, perform flash_attn instead
|
||||
if gate_mask.sum() == 0:
|
||||
return flash_attn_varlen_func(
|
||||
q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen, causal=False
|
||||
)
|
||||
|
||||
# Determine which query positions are selected.
|
||||
# nonzero_indices has shape [N, 3] where each row is [chunk_index, head_index, seq_index].
|
||||
moba_q_indices = gate_mask.reshape(gate_mask.shape[0], -1).nonzero(as_tuple=True)[-1] # [(h s k)]
|
||||
moba_q_sh_indices = (moba_q_indices % seqlen) * num_head + (moba_q_indices // seqlen)
|
||||
moba_q = rearrange(q, "s h d -> (h s) d").index_select(0, moba_q_indices).unsqueeze(1)
|
||||
|
||||
# Build cumulative sequence lengths for the selected queries.
|
||||
moba_seqlen_q = gate_mask.sum(dim=-1).flatten()
|
||||
q_zero_mask = moba_seqlen_q == 0
|
||||
valid_expert_mask = ~q_zero_mask
|
||||
if q_zero_mask.sum() > 0:
|
||||
moba_seqlen_q = moba_seqlen_q[valid_expert_mask]
|
||||
moba_cu_seqlen_q = torch.cat(
|
||||
(
|
||||
torch.tensor([0], device=q.device, dtype=moba_seqlen_q.dtype),
|
||||
moba_seqlen_q.cumsum(dim=0),
|
||||
),
|
||||
dim=0,
|
||||
).to(torch.int32)
|
||||
|
||||
# Rearrange gathered KV for the MOBA branch.
|
||||
experts_tensor = rearrange(gathered, "nc cl two h d -> (nc h) cl two d")
|
||||
valid_expert_lengths = chunk_lengths.unsqueeze(1).expand(num_filtered_chunk, num_head).reshape(-1).to(torch.int32)
|
||||
if q_zero_mask.sum() > 0:
|
||||
experts_tensor = experts_tensor[valid_expert_mask]
|
||||
valid_expert_lengths = valid_expert_lengths[valid_expert_mask]
|
||||
|
||||
seq_range = torch.arange(experts_tensor.shape[1], device=experts_tensor.device).unsqueeze(0)
|
||||
mask = seq_range < valid_expert_lengths.unsqueeze(1)
|
||||
moba_kv = experts_tensor[mask] # Shape: ((nc h cl_valid) two d)
|
||||
moba_kv = moba_kv.unsqueeze(2) # Shape: ((nc h cl_valid) two 1 d)
|
||||
|
||||
moba_cu_seqlen_kv = torch.cat(
|
||||
[torch.zeros(1, device=experts_tensor.device, dtype=torch.int32),
|
||||
valid_expert_lengths.cumsum(dim=0)],
|
||||
dim=0,
|
||||
).to(torch.int32)
|
||||
|
||||
assert (
|
||||
moba_cu_seqlen_kv.shape == moba_cu_seqlen_q.shape
|
||||
), f"Mismatch between moba_cu_seqlen_kv.shape and moba_cu_seqlen_q.shape: {moba_cu_seqlen_kv.shape} vs {moba_cu_seqlen_q.shape}"
|
||||
|
||||
return MixedAttention.apply(
|
||||
q,
|
||||
k,
|
||||
v,
|
||||
self_attn_cu_seqlen,
|
||||
moba_q,
|
||||
moba_kv,
|
||||
moba_cu_seqlen_q,
|
||||
moba_cu_seqlen_kv,
|
||||
max_seqlen,
|
||||
moba_chunk_size,
|
||||
moba_q_sh_indices,
|
||||
)
|
||||
|
||||
|
||||
def process_moba_input(
|
||||
x,
|
||||
patch_resolution,
|
||||
chunk_size,
|
||||
):
|
||||
"""
|
||||
Process inputs for the attention function.
|
||||
|
||||
Args:
|
||||
x (torch.Tensor): Input tensor with shape [batch_size, num_patches, num_heads, head_dim].
|
||||
patch_resolution (tuple): Tuple containing the patch resolution (t, h, w).
|
||||
chunk_size (int): Size of the chunk. (maybe tuple or int, according to chunk type)
|
||||
|
||||
Returns:
|
||||
torch.Tensor: Processed input tensor.
|
||||
"""
|
||||
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
|
||||
moba_chunk_size = int(chunk_size * patch_resolution[1] * patch_resolution[2])
|
||||
else:
|
||||
assert isinstance(chunk_size, (Tuple, list)), f"chunk_size should be a tuple, list, or int, now it is: {type(chunk_size)}"
|
||||
if len(chunk_size) == 2:
|
||||
assert patch_resolution[1] % chunk_size[0] == 0 and patch_resolution[2] % chunk_size[1] == 0, f"spatial patch_resolution {patch_resolution[1:]} should be divisible by 2d chunk_size {chunk_size}"
|
||||
nch, ncw = patch_resolution[1] // chunk_size[0], patch_resolution[2] // chunk_size[1]
|
||||
x = rearrange(x, "b (t nch ch ncw cw) n d -> b (nch ncw t ch cw) n d", t=patch_resolution[0], nch=nch, ncw=ncw, ch=chunk_size[0], cw=chunk_size[1])
|
||||
moba_chunk_size = patch_resolution[0] * chunk_size[0] * chunk_size[1]
|
||||
elif len(chunk_size) == 3:
|
||||
assert patch_resolution[0] % chunk_size[0] == 0 and patch_resolution[1] % chunk_size[1] == 0 and patch_resolution[2] % chunk_size[2] == 0, f"patch_resolution {patch_resolution} should be divisible by 3d chunk_size {chunk_size}"
|
||||
nct, nch, ncw = patch_resolution[0] // chunk_size[0], patch_resolution[1] // chunk_size[1], patch_resolution[2] // chunk_size[2]
|
||||
x = rearrange(x, "b (nct ct nch ch ncw cw) n d -> b (nct nch ncw ct ch cw) n d", nct=nct, nch=nch, ncw=ncw, ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
|
||||
moba_chunk_size = chunk_size[0] * chunk_size[1] * chunk_size[2]
|
||||
else:
|
||||
raise ValueError(f"chunk_size should be a int, or a tuple of length 2 or 3, now it is: {len(chunk_size)}")
|
||||
|
||||
return x, moba_chunk_size
|
||||
|
||||
|
||||
def process_moba_output(
|
||||
x,
|
||||
patch_resolution,
|
||||
chunk_size,
|
||||
):
|
||||
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
|
||||
pass
|
||||
elif len(chunk_size) == 2:
|
||||
x = rearrange(x, "b (nch ncw t ch cw) n d -> b (t nch ch ncw cw) n d", nch=patch_resolution[1] // chunk_size[0], ncw=patch_resolution[2] // chunk_size[1], t=patch_resolution[0], ch=chunk_size[0], cw=chunk_size[1])
|
||||
elif len(chunk_size) == 3:
|
||||
x = rearrange(x, "b (nct nch ncw ct ch cw) n d -> b (nct ct nch ch ncw cw) n d", nct=patch_resolution[0] // chunk_size[0], nch=patch_resolution[1] // chunk_size[1], ncw=patch_resolution[2] // chunk_size[2], ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
|
||||
|
||||
return x
|
||||
|
||||
|
||||
# TEST
|
||||
def generate_data(batch_size, seqlen, num_head, head_dim, dtype):
|
||||
random.seed(0)
|
||||
torch.manual_seed(0)
|
||||
torch.cuda.manual_seed(0)
|
||||
device = torch.cuda.current_device()
|
||||
|
||||
q = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
k = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
v = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
|
||||
print(f"q.shape: {q.shape}, k.shape: {k.shape}, v.shape: {v.shape}")
|
||||
cu_seqlens = torch.arange(0, q.shape[0] * q.shape[1] + 1, q.shape[1], dtype=torch.int32, device='cuda')
|
||||
max_seqlen = q.shape[1]
|
||||
q = rearrange(q, "b s ... -> (b s) ...")
|
||||
k = rearrange(k, "b s ... -> (b s) ...")
|
||||
v = rearrange(v, "b s ... -> (b s) ...")
|
||||
|
||||
return q, k, v, cu_seqlens, max_seqlen
|
||||
|
||||
|
||||
def test_attn_varlen_moba_speed(batch, head, seqlen, head_dim, moba_chunk_size, moba_topk, dtype=torch.bfloat16, select_mode='threshold', simsum_threshold=0.25, threshold_type='query_head'):
|
||||
"""Speed test comparing flash_attn vs moba_attention"""
|
||||
# Get data
|
||||
q, k, v, cu_seqlen, max_seqlen = generate_data(batch, seqlen, head, head_dim, dtype)
|
||||
print(f"batch:{batch} head:{head} seqlen:{seqlen} chunk:{moba_chunk_size} topk:{moba_topk} select_mode: {select_mode} simsum_threshold:{simsum_threshold}")
|
||||
vo_grad = torch.randn_like(q)
|
||||
|
||||
# Warmup
|
||||
warmup_iters = 3
|
||||
perf_test_iters = 10
|
||||
|
||||
# Warmup
|
||||
for _ in range(warmup_iters):
|
||||
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
|
||||
torch.autograd.backward(o, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
start_flash = time.perf_counter()
|
||||
for _ in range(perf_test_iters):
|
||||
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
|
||||
torch.autograd.backward(o, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
time_flash = (time.perf_counter() - start_flash) / perf_test_iters * 1000
|
||||
|
||||
# Warmup
|
||||
for _ in range(warmup_iters):
|
||||
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
|
||||
torch.autograd.backward(om, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
start_moba = time.perf_counter()
|
||||
for _ in range(perf_test_iters):
|
||||
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
|
||||
torch.autograd.backward(om, vo_grad)
|
||||
|
||||
torch.cuda.synchronize()
|
||||
time_moba = (time.perf_counter() - start_moba) / perf_test_iters * 1000
|
||||
|
||||
print(f"Flash: {time_flash:.2f}ms, MoBA: {time_moba:.2f}ms")
|
||||
print(f"Speedup: {time_flash / time_moba:.2f}x")
|
||||
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
"""
|
||||
CUDA_VISIBLE_DEVICES=1 \
|
||||
python -u csrc/attn/vmoba_attn/vmoba/vmoba.py
|
||||
"""
|
||||
test_attn_varlen_moba_speed(batch=1, head=12, seqlen=32760, head_dim=128, moba_chunk_size=32760 // 3 // 6 // 4, moba_topk=3, select_mode='threshold', simsum_threshold=0.3, threshold_type='query_head')
|
||||
@@ -1,72 +0,0 @@
|
||||
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
wget \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
zsh \
|
||||
vim \
|
||||
curl \
|
||||
gcc-11 \
|
||||
g++-11 \
|
||||
clang-11 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Set up C++20 compilers for ThunderKittens
|
||||
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Set CUDA environment variables
|
||||
ENV CUDA_HOME=/usr/local/cuda-12.8
|
||||
ENV PATH=${CUDA_HOME}/bin:${PATH}
|
||||
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Install uv and source its environment
|
||||
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
|
||||
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject.toml ./
|
||||
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
# Create and activate virtual environment with specific Python version and seed
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.10 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
@@ -1,72 +0,0 @@
|
||||
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
wget \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
zsh \
|
||||
vim \
|
||||
curl \
|
||||
gcc-11 \
|
||||
g++-11 \
|
||||
clang-11 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Set up C++20 compilers for ThunderKittens
|
||||
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Set CUDA environment variables
|
||||
ENV CUDA_HOME=/usr/local/cuda-12.8
|
||||
ENV PATH=${CUDA_HOME}/bin:${PATH}
|
||||
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Install uv and source its environment
|
||||
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
|
||||
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject.toml ./
|
||||
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
# Create and activate virtual environment with specific Python version and seed
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.11 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
@@ -1,72 +0,0 @@
|
||||
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
wget \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
zsh \
|
||||
vim \
|
||||
curl \
|
||||
gcc-11 \
|
||||
g++-11 \
|
||||
clang-11 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Set up C++20 compilers for ThunderKittens
|
||||
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Set CUDA environment variables
|
||||
ENV CUDA_HOME=/usr/local/cuda-12.8
|
||||
ENV PATH=${CUDA_HOME}/bin:${PATH}
|
||||
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Install uv and source its environment
|
||||
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
|
||||
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject.toml ./
|
||||
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
# Create and activate virtual environment with specific Python version and seed
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.12 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
@@ -1,72 +0,0 @@
|
||||
FROM nvidia/cuda:12.9.1-cudnn-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
wget \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
zsh \
|
||||
vim \
|
||||
curl \
|
||||
gcc-11 \
|
||||
g++-11 \
|
||||
clang-11 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Set up C++20 compilers for ThunderKittens
|
||||
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Set CUDA environment variables
|
||||
ENV CUDA_HOME=/usr/local/cuda-12.9
|
||||
ENV PATH=${CUDA_HOME}/bin:${PATH}
|
||||
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Install uv and source its environment
|
||||
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
|
||||
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject.toml ./
|
||||
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
# Create and activate virtual environment with specific Python version and seed
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.12 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
@@ -1,26 +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"
|
||||
rm -rf "$(SOURCEDIR)/inference/examples"
|
||||
rm -rf "$(SOURCEDIR)/training/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.
|
||||
@@ -0,0 +1,68 @@
|
||||
|
||||
|
||||
|
||||
## 🧱 Data Preprocess
|
||||
|
||||
To save GPU memory, we precompute text embeddings and VAE latents to eliminate the need to load the text encoder and VAE during training.
|
||||
|
||||
|
||||
We provide a sample dataset to help you get started. Download the source media using the following command:
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Image-Vid-Finetune-Src --local_dir=data/Image-Vid-Finetune-Src --repo_type=dataset
|
||||
```
|
||||
To preprocess the dataset for fine-tuning or distillation, run:
|
||||
```
|
||||
bash scripts/preprocess/preprocess_mochi_data.sh # for mochi
|
||||
bash scripts/preprocess/preprocess_hunyuan_data.sh # for hunyuan
|
||||
```
|
||||
|
||||
The preprocessed dataset will be stored in `Image-Vid-Finetune-Mochi` or `Image-Vid-Finetune-HunYuan` correspondingly.
|
||||
|
||||
### Process your own dataset
|
||||
|
||||
If you wish to create your own dataset for finetuning or distillation, please structure you video dataset in the following format:
|
||||
|
||||
path_to_dataset_folder/
|
||||
├── media/
|
||||
│ ├── 0.jpg
|
||||
│ ├── 1.mp4
|
||||
│ ├── 2.jpg
|
||||
├── video2caption.json
|
||||
└── merge.txt
|
||||
|
||||
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,
|
||||
```
|
||||
{
|
||||
"path": "1.mp4",
|
||||
"resolution": {
|
||||
"width": 848,
|
||||
"height": 480
|
||||
},
|
||||
"fps": 30.0,
|
||||
"duration": 6.033333333333333,
|
||||
"cap": [
|
||||
"caption"
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
Use a txt file (merge.txt) to contain the source folder for media and the JSON file for meta information:
|
||||
|
||||
```
|
||||
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,15 +0,0 @@
|
||||
sphinx==7.4.7
|
||||
sphinx-argparse==0.5.2
|
||||
sphinx-autodoc2==0.5.0
|
||||
sphinx-book-theme==1.1.4
|
||||
sphinx-copybutton==0.5.2
|
||||
sphinx-design==0.6.1
|
||||
sphinx-togglebutton==0.3.2
|
||||
myst-parser==3.0.1
|
||||
msgspec
|
||||
commonmark # Required by sphinx-argparse when using :markdownhelp:
|
||||
|
||||
# packages to install to build the documentation
|
||||
cachetools
|
||||
# -f https://download.pytorch.org/whl/cpu
|
||||
torch
|
||||
@@ -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));
|
||||
});
|
||||
});
|
||||
|
Before Width: | Height: | Size: 194 KiB |
|
Before Width: | Height: | Size: 303 KiB |
|
Before Width: | Height: | Size: 18 KiB |
|
Before Width: | Height: | Size: 27 KiB |
|
Before Width: | Height: | Size: 40 KiB |
@@ -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,19 +0,0 @@
|
||||
# Summary
|
||||
|
||||
## Video Generator
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.VideoGenerator
|
||||
```
|
||||
|
||||
## Initialization Configuration
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.configs.pipelines.PipelineConfig
|
||||
```
|
||||
|
||||
## Sampling Configuration
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.configs.sample.SamplingParam
|
||||
```
|
||||
@@ -1,22 +0,0 @@
|
||||
# type: ignore
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from docutils import nodes
|
||||
from myst_parser.parsers.sphinx_ import MystParser
|
||||
from sphinx.ext.napoleon import docstring
|
||||
|
||||
|
||||
class NapoleonParser(MystParser):
|
||||
|
||||
def parse(self, input_string: str, document: nodes.document) -> None:
|
||||
# Get the Sphinx configuration
|
||||
config = document.settings.env.config
|
||||
|
||||
parsed_content = str(
|
||||
docstring.GoogleDocstring(
|
||||
str(docstring.NumpyDocstring(input_string, config)),
|
||||
config,
|
||||
))
|
||||
return super().parse(parsed_content, document)
|
||||
|
||||
|
||||
Parser = NapoleonParser
|
||||
@@ -1,275 +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 logging
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import requests
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
REPO_ROOT = Path(__file__).resolve().parent.parent.parent
|
||||
print(os.path.abspath(REPO_ROOT))
|
||||
sys.path.append(os.path.abspath(REPO_ROOT))
|
||||
|
||||
# -- 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",
|
||||
"autodoc2",
|
||||
"myst_parser",
|
||||
"sphinxarg.ext",
|
||||
"sphinx_design",
|
||||
"sphinx_togglebutton",
|
||||
]
|
||||
myst_enable_extensions = [
|
||||
"colon_fence",
|
||||
"fieldlist",
|
||||
]
|
||||
autodoc2_packages = [
|
||||
{
|
||||
"path": "../../fastvideo",
|
||||
"exclude_dirs": ["__pycache__", "third_party"],
|
||||
},
|
||||
]
|
||||
autodoc2_output_dir = "api"
|
||||
autodoc2_render_plugin = "myst"
|
||||
autodoc2_hidden_objects = ["dunder", "private", "inherited"]
|
||||
autodoc2_docstring_parser_regexes = [
|
||||
(".*", "docs.source.autodoc2_docstring_parser"),
|
||||
]
|
||||
autodoc2_sort_names = True
|
||||
autodoc2_index_template = None
|
||||
autodoc2_skip_module_regexes = [
|
||||
"fastvideo.dataset",
|
||||
"fastvideo.distill",
|
||||
"fastvideo.data_preprocess",
|
||||
"fastvideo.models",
|
||||
"fastvideo.sample",
|
||||
"fastvideo.utils",
|
||||
"fastvideo.distill_adv",
|
||||
"fastvideo.train",
|
||||
]
|
||||
|
||||
# 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/logos/icon_simple.svg'
|
||||
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,
|
||||
# Prevents the full API being added to the left sidebar of every page.
|
||||
# Reduces build time by 2.5x and reduces build size from ~225MB to ~95MB.
|
||||
'collapse_navbar': True,
|
||||
# Makes API visible in the right sidebar on API reference pages.
|
||||
'show_toc_level': 3,
|
||||
}
|
||||
# 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[str | None, str | None]:
|
||||
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
|
||||
|
||||
# Get path from module name
|
||||
file = Path(f"{info['module'].replace('.', '/')}.py")
|
||||
path = REPO_ROOT / file
|
||||
if not path.exists():
|
||||
path = REPO_ROOT / file.with_suffix("") / "__init__.py"
|
||||
if not path.exists():
|
||||
return None
|
||||
|
||||
# Get the line number of the object
|
||||
with open(path) as f:
|
||||
lines = f.readlines()
|
||||
name = info['fullname'].split(".")[-1]
|
||||
pattern = fr"^( {{4}})*((def|class) )?{name}\b.*"
|
||||
for lineno, line in enumerate(lines, 1):
|
||||
if not line or line.startswith("#"):
|
||||
continue
|
||||
if re.match(pattern, line):
|
||||
break
|
||||
|
||||
# If the line number is not found, return None
|
||||
if lineno == len(lines):
|
||||
return None
|
||||
|
||||
# If the line number is found, create the URL
|
||||
filename = path.relative_to(REPO_ROOT)
|
||||
if "checkouts" in path.parts:
|
||||
# a PR build on readthedocs
|
||||
pr_number = REPO_ROOT.name
|
||||
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",
|
||||
"huggingface_hub",
|
||||
"torchvision",
|
||||
"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)
|
||||
|
||||
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),
|
||||
}
|
||||
|
||||
navigation_with_keys = False
|
||||
@@ -1,32 +0,0 @@
|
||||
(docker)=
|
||||
# 🐳 Using the FastVideo Docker Image
|
||||
|
||||
If you prefer a containerized development environment or want to avoid managing dependencies manually, you can use our prebuilt Docker image:
|
||||
|
||||
**Images:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
|
||||
|
||||
## Starting the container
|
||||
|
||||
```bash
|
||||
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest
|
||||
```
|
||||
|
||||
This will:
|
||||
|
||||
- Start the container with GPU access
|
||||
- Drop you into a shell with the `fastvideo-dev` Conda environment preconfigured
|
||||
|
||||
## Using the container
|
||||
|
||||
```bash
|
||||
# Conda environment should already be active
|
||||
# FastVideo package installed in editable mode
|
||||
|
||||
# Pull the latest changes from remote
|
||||
cd /FastVideo
|
||||
git pull
|
||||
|
||||
# Run linters and tests
|
||||
pre-commit run --all-files
|
||||
pytest tests/
|
||||
```
|
||||
@@ -1,13 +0,0 @@
|
||||
(developer-env)
|
||||
|
||||
# 🧰 Developer Environment
|
||||
|
||||
Accelerate your FastVideo development workflow by leveraging Docker images and cloud GPUs for efficient experimentation and reproducible environments.
|
||||
|
||||
:::{toctree}
|
||||
:caption: Contents
|
||||
:maxdepth: 1
|
||||
|
||||
docker
|
||||
runpod
|
||||
:::
|
||||
@@ -1,54 +0,0 @@
|
||||
(runpod)=
|
||||
|
||||
# 📦 Developing FastVideo on RunPod
|
||||
|
||||
You can easily use the FastVideo Docker image as a custom container on [RunPod](https://www.runpod.io) for development or experimentation.
|
||||
|
||||
## Creating a new pod
|
||||
|
||||
Choose a GPU that supports CUDA 12.8
|
||||
|
||||
Pick 1 or 2 L40S GPU(s)
|
||||
|
||||

|
||||
|
||||
When creating your pod template, use this image:
|
||||
|
||||
```
|
||||
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest
|
||||
```
|
||||
|
||||
Paste Container Start Command to support SSH ([RunPod Docs](https://docs.runpod.io/pods/configuration/use-ssh)):
|
||||
|
||||
```bash
|
||||
bash -c "apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-server -y;mkdir -p ~/.ssh;cd $_;chmod 700 ~/.ssh;echo \"$PUBLIC_KEY\" >> authorized_keys;chmod 700 authorized_keys;service ssh start;sleep infinity"
|
||||
```
|
||||
|
||||

|
||||
|
||||
After deploying, the pod will take a few minutes to pull the image and start the SSH service.
|
||||
|
||||

|
||||
|
||||
## Working with the pod
|
||||
|
||||
After SSH'ing into your pod, you'll find the `fastvideo-dev` Conda environment already activated.
|
||||
|
||||
To pull in the latest changes from the GitHub repo:
|
||||
|
||||
```bash
|
||||
cd /FastVideo
|
||||
git pull
|
||||
```
|
||||
|
||||
`If you have a persistent volume and want to keep your code changes, you can move /FastVideo to /workspace/FastVideo, or simply clone the repository there.`
|
||||
|
||||
Run your development workflows as usual:
|
||||
|
||||
```bash
|
||||
# Run linters
|
||||
pre-commit run --all-files
|
||||
|
||||
# Run tests
|
||||
pytest tests/
|
||||
```
|
||||
@@ -1,73 +0,0 @@
|
||||
(developer-overview)=
|
||||
|
||||
# 🛠️ 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.12 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
Install `uv` (optional, but recommended):
|
||||
|
||||
From instructions on [uv](https://astral.sh/uv/):
|
||||
|
||||
```
|
||||
curl -LsSf https://astral.sh/uv/install.sh | sh
|
||||
# or
|
||||
wget -qO- https://astral.sh/uv/install.sh | sh
|
||||
```
|
||||
|
||||
Clone the FastVideo repository and go to the FastVideo directory:
|
||||
|
||||
```
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
|
||||
```
|
||||
|
||||
Now you can install FastVideo and setup git hooks for running linting. By using `pre-commit`, the linters will run and have to pass before you'll be able to make a commit.
|
||||
|
||||
```bash
|
||||
uv pip install -e .[dev]
|
||||
|
||||
# Can also install flash-attn (optional)
|
||||
uv pip install flash-attn --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/
|
||||
```
|
||||
|
||||
If you are on a Hopper GPU, you should also install [FA3](https://github.com/Dao-AILab/flash-attention) for much better performance:
|
||||
|
||||
```
|
||||
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention/hopper
|
||||
|
||||
# make sure you have ninja installed
|
||||
uv pip install ninja
|
||||
|
||||
python setup.py install
|
||||
```
|
||||