Compare commits
120
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8573d4f05e | ||
|
|
210a733515 | ||
|
|
0aef0e6f63 | ||
|
|
dd022ad9be | ||
|
|
832ad61e5b | ||
|
|
9419c04ee3 | ||
|
|
a37b39d83c | ||
|
|
bb8c769c8e | ||
|
|
576c214f28 | ||
|
|
b79d1fc15b | ||
|
|
eb66e1c18d | ||
|
|
616d43c1cf | ||
|
|
7244a4b27f | ||
|
|
7e5ebb4582 | ||
|
|
65ed588570 | ||
|
|
14adfe2edc | ||
|
|
6198c6a640 | ||
|
|
e6b71b531b | ||
|
|
ae1d112c6a | ||
|
|
bf4de1f38f | ||
|
|
66fdcc8e76 | ||
|
|
ed1e8d6bad | ||
|
|
b9423ca3f8 | ||
|
|
ad16289871 | ||
|
|
2a41da1e6b | ||
|
|
32133171da | ||
|
|
19674c6f29 | ||
|
|
508afb7002 | ||
|
|
288ea88105 | ||
|
|
eb0f1318f3 | ||
|
|
ce9b5910cc | ||
|
|
d0e5a6214a | ||
|
|
834562b2db | ||
|
|
060cc7b9ba | ||
|
|
6c58a5ba62 | ||
|
|
48d9f61f86 | ||
|
|
5f938b5844 | ||
|
|
74da2a7370 | ||
|
|
580d6dfe1f | ||
|
|
344e43006a | ||
|
|
c5155b256e | ||
|
|
e005c7f3ac | ||
|
|
ff5a79ef60 | ||
|
|
ab01dc4ba5 | ||
|
|
285a950c1b | ||
|
|
46a0a85d85 | ||
|
|
4aeabbc629 | ||
|
|
949bb5c835 | ||
|
|
aab74c1271 | ||
|
|
f89d86944f | ||
|
|
8741d204a5 | ||
|
|
cdc85f58a8 | ||
|
|
0262d2f089 | ||
|
|
62c0343465 | ||
|
|
1e1a023fb0 | ||
|
|
1d2517ad8e | ||
|
|
d41186cb4a | ||
|
|
78e0c7eec9 | ||
|
|
1c41a94b62 | ||
|
|
2e66aafe20 | ||
|
|
55074bda76 | ||
|
|
de65bec2b7 | ||
|
|
7664dd0de3 | ||
|
|
019a88ced4 | ||
|
|
72de11abcc | ||
|
|
d71a4ebffc | ||
|
|
1089ab43bf | ||
|
|
97d4b984c9 | ||
|
|
2a8953d74d | ||
|
|
8801b10da7 | ||
|
|
6b413f2ec4 | ||
|
|
28b72694aa | ||
|
|
4afb0cfe4f | ||
|
|
3eec1281cf | ||
|
|
0660489e38 | ||
|
|
dd871a17bf | ||
|
|
dc11529862 | ||
|
|
ffabf85e31 | ||
|
|
c0026ca5ba | ||
|
|
0f2bbe71ac | ||
|
|
2a46902ecb | ||
|
|
66012d3a4c | ||
|
|
f666b9de41 | ||
|
|
7e3c073b55 | ||
|
|
a6aa21bd07 | ||
|
|
6519b57aab | ||
|
|
675aea6ece | ||
|
|
46e7a15e0d | ||
|
|
e4f702d7ec | ||
|
|
bb68fcc809 | ||
|
|
b392e6a874 | ||
|
|
0991003905 | ||
|
|
8f8ce6d9e1 | ||
|
|
e3d0cbe185 | ||
|
|
d5ec468d43 | ||
|
|
e55fa6e5dc | ||
|
|
61b6ddeee1 | ||
|
|
a9a000f45d | ||
|
|
66b8b8561e | ||
|
|
6684872616 | ||
|
|
7f654e3332 | ||
|
|
8631c1b806 | ||
|
|
5357e12b5a | ||
|
|
bdfdf1dfee | ||
|
|
d156461785 | ||
|
|
6edf113838 | ||
|
|
dcf7738cbc | ||
|
|
b2ebaaf865 | ||
|
|
7768bb80f6 | ||
|
|
a335811869 | ||
|
|
357b0533fe | ||
|
|
2ec3732758 | ||
|
|
a004408a93 | ||
|
|
007e237e69 | ||
|
|
8e18dc9f71 | ||
|
|
7ab32539af | ||
|
|
6ef8fcb61d | ||
|
|
016e24da63 | ||
|
|
85b8717545 | ||
|
|
657fd745e1 |
@@ -0,0 +1,166 @@
|
||||
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/v1/models/encoders/**"
|
||||
- "fastvideo/v1/models/loader/**"
|
||||
- "fastvideo/v1/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/v1/models/vaes/**"
|
||||
- "fastvideo/v1/models/loader/**"
|
||||
- "fastvideo/v1/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/v1/models/dits/**"
|
||||
- "fastvideo/v1/models/loader/**"
|
||||
- "fastvideo/v1/tests/transformers/**"
|
||||
- "fastvideo/v1/layers/**"
|
||||
- "fastvideo/v1/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/v1/**/*.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/v1/tests/lora/**"
|
||||
- "fastvideo/v1/models/loader/**"
|
||||
- "fastvideo/v1/tests/transformers/**"
|
||||
- "fastvideo/v1/pipelines/**"
|
||||
- "fastvideo/v1/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/v1/**"
|
||||
- "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/v1/**"
|
||||
- "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/v1/**"
|
||||
- "csrc/attn/vsa/**"
|
||||
- "csrc/attn/tk/**"
|
||||
- "csrc/attn/setup_vsa.py"
|
||||
- "csrc/attn/config_vsa.py"
|
||||
- "csrc/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/v1/**"
|
||||
- "csrc/attn/st_attn/**"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/attn/config_sta.py"
|
||||
- "csrc/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/st_attn/**"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/attn/config_sta.py"
|
||||
- "csrc/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/vsa/**"
|
||||
- "csrc/attn/tk/**"
|
||||
- "csrc/attn/setup_vsa.py"
|
||||
- "csrc/attn/config_vsa.py"
|
||||
- "csrc/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"
|
||||
Executable
+125
@@ -0,0 +1,125 @@
|
||||
#!/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/v1/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"
|
||||
;;
|
||||
*)
|
||||
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
|
||||
@@ -0,0 +1,40 @@
|
||||
#!/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,14 +4,6 @@ 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/collect_env.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
|
||||
@@ -25,5 +17,13 @@ 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 fastvideo/utils/collect_env.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
@@ -0,0 +1,56 @@
|
||||
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
|
||||
@@ -160,8 +160,7 @@ def execute_command(pod_id):
|
||||
setup_steps = [
|
||||
"tar -xzf /tmp/repo.tar.gz --no-same-owner -C /workspace/",
|
||||
f"cd /workspace/{repo_name}",
|
||||
"source /opt/conda/etc/profile.d/conda.sh",
|
||||
"conda activate fastvideo-dev",
|
||||
"source $HOME/.local/bin/env && source /opt/venv/bin/activate",
|
||||
args.test_command
|
||||
]
|
||||
|
||||
|
||||
@@ -13,4 +13,4 @@
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
+212
-20
@@ -12,13 +12,11 @@ on:
|
||||
paths:
|
||||
- "fastvideo/**/*.py"
|
||||
- ".github/workflows/pr-test.yml"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
- "csrc/**"
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
custom_image:
|
||||
description: "Custom image from this repository (default: fastvideo-dev:latest)"
|
||||
required: false
|
||||
default: "fastvideo-dev:latest"
|
||||
type: string
|
||||
run_encoder_test:
|
||||
description: "Run encoder-test"
|
||||
required: false
|
||||
@@ -39,10 +37,41 @@ on:
|
||||
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
|
||||
@@ -59,24 +88,72 @@ jobs:
|
||||
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/st_attn/**'
|
||||
- 'csrc/attn/setup_sta.py'
|
||||
- 'csrc/attn/config_sta.py'
|
||||
- 'csrc/attn/st_attn.cpp'
|
||||
vsa-kernel-paths: &vsa-kernel-paths
|
||||
- 'csrc/attn/vsa/**'
|
||||
- 'csrc/attn/tk/**'
|
||||
- 'csrc/attn/setup_vsa.py'
|
||||
- 'csrc/attn/config_vsa.py'
|
||||
- 'csrc/attn/vsa.cpp'
|
||||
vsa-paths: &vsa-paths
|
||||
- 'fastvideo/v1/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
|
||||
# Actual tests
|
||||
encoder-test:
|
||||
- 'fastvideo/v1/models/encoders/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/encoders/**'
|
||||
- *common-paths
|
||||
vae-test:
|
||||
- 'fastvideo/v1/models/vaes/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/vaes/**'
|
||||
- *common-paths
|
||||
transformer-test:
|
||||
- 'fastvideo/v1/models/dits/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/transformers/**'
|
||||
- 'fastvideo/v1/layers/**'
|
||||
- 'fastvideo/v1/attention/**'
|
||||
- *common-paths
|
||||
training-test:
|
||||
- 'fastvideo/v1/**'
|
||||
- *common-paths
|
||||
training-test-VSA:
|
||||
- 'fastvideo/v1/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
inference-test-STA:
|
||||
- 'fastvideo/v1/**'
|
||||
- *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
|
||||
@@ -89,8 +166,8 @@ jobs:
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -107,8 +184,8 @@ jobs:
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -125,8 +202,8 @@ jobs:
|
||||
gpu_type: "NVIDIA L40S"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -135,8 +212,7 @@ jobs:
|
||||
ssim-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
|
||||
github.event_name != 'workflow_dispatch' || (github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
@@ -153,14 +229,130 @@ jobs:
|
||||
volume_size: 200
|
||||
disk_size: 200
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:${{ matrix.python-version.tag }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/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/v1/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/v1/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/v1/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_block_sparse.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/v1/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:
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
|
||||
# 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:
|
||||
@@ -177,7 +369,7 @@ jobs:
|
||||
|
||||
- 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"]'
|
||||
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
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
@@ -10,7 +10,7 @@ jobs:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
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
|
||||
|
||||
@@ -43,6 +43,8 @@ on:
|
||||
required: true
|
||||
RUNPOD_PRIVATE_KEY:
|
||||
required: true
|
||||
WANDB_API_KEY:
|
||||
required: false
|
||||
|
||||
jobs:
|
||||
run-test:
|
||||
@@ -55,7 +57,7 @@ jobs:
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Set up SSH key
|
||||
run: |
|
||||
@@ -72,6 +74,7 @@ jobs:
|
||||
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
|
||||
|
||||
@@ -5,7 +5,7 @@ on:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/sliding_tile_attention/setup.py"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
@@ -23,13 +23,13 @@ jobs:
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/sliding_tile_attention
|
||||
cd csrc/attn
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_sta.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")
|
||||
OLD_VERSION=$(git show HEAD~1:./setup_sta.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
@@ -136,19 +136,21 @@ 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/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
|
||||
cd csrc/attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_sta.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/sliding_tile_attention
|
||||
cd csrc/attn
|
||||
|
||||
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)
|
||||
@@ -163,7 +165,7 @@ jobs:
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}
|
||||
path: csrc/sliding_tile_attention/dist/*.whl
|
||||
path: csrc/attn/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
@@ -229,17 +231,19 @@ 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/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
|
||||
cd csrc/attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_sta.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/sliding_tile_attention/dist/
|
||||
packages-dir: csrc/attn/dist/
|
||||
|
||||
@@ -28,4 +28,4 @@ jobs:
|
||||
|
||||
- name: Run Pytest
|
||||
run: |
|
||||
pytest --ignore csrc/sliding_tile_attention/test
|
||||
pytest --ignore csrc/attn/test
|
||||
|
||||
+4
-1
@@ -20,6 +20,7 @@ samples/
|
||||
data/
|
||||
outputs/
|
||||
outputs_video
|
||||
checkpoints/
|
||||
sbatch.sh
|
||||
*.out
|
||||
env
|
||||
@@ -27,7 +28,6 @@ env
|
||||
**/build/
|
||||
**.pyc
|
||||
**.txt
|
||||
**.json
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
@@ -41,6 +41,7 @@ eggs/
|
||||
docs/_build/
|
||||
docs/source/getting_started/examples/
|
||||
docs/source/inference/examples/
|
||||
docs/source/training/examples/
|
||||
|
||||
# VSCode
|
||||
.vscode/
|
||||
@@ -60,3 +61,5 @@ docs/source/inference/examples/
|
||||
|
||||
# Static images
|
||||
!docs/source/_static/images/**/*.png
|
||||
!comfyui/assets/**/*.png
|
||||
!comfyui/assets/**/*.gif
|
||||
|
||||
+2
-2
@@ -1,3 +1,3 @@
|
||||
[submodule "csrc/sliding_tile_attention/tk"]
|
||||
path = csrc/sliding_tile_attention/tk
|
||||
[submodule "csrc/attn/tk"]
|
||||
path = csrc/attn/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
@@ -33,7 +33,7 @@ repos:
|
||||
args: [--in-place, --verbose]
|
||||
additional_dependencies: [toml] # TODO: Remove when yapf is upgraded
|
||||
- repo: https://github.com/astral-sh/ruff-pre-commit
|
||||
rev: v0.11.4
|
||||
rev: v0.11.12
|
||||
hooks:
|
||||
- id: ruff
|
||||
args: [--output-format, github, --fix]
|
||||
@@ -48,7 +48,7 @@ repos:
|
||||
hooks:
|
||||
- id: isort
|
||||
- repo: https://github.com/jackdewinter/pymarkdown
|
||||
rev: v0.9.29
|
||||
rev: v0.9.30
|
||||
hooks:
|
||||
- id: pymarkdown
|
||||
args: [fix]
|
||||
@@ -60,7 +60,7 @@ repos:
|
||||
rev: v1.15.0
|
||||
hooks:
|
||||
- id: mypy
|
||||
args: [--python-version, '3.10', --follow-imports, "skip", ]
|
||||
args: [--python-version, '3.10', --follow-imports, "skip" ]
|
||||
additional_dependencies: [types-cachetools, types-setuptools, types-PyYAML, types-requests]
|
||||
- repo: local
|
||||
hooks:
|
||||
@@ -69,7 +69,7 @@ repos:
|
||||
entry: bash
|
||||
args:
|
||||
- -c
|
||||
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
|
||||
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/" | grep -v "^fastvideo/v1/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
|
||||
|
||||
@@ -8,13 +8,18 @@ It features a clean, consistent API that works across popular video models, maki
|
||||
With FastVideo's optimizations, you can achieve more than 3x inference improvement compared to other systems.
|
||||
|
||||
<p align="center">
|
||||
| <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/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> <b>Slack</b> </a> |
|
||||
| <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/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ" target="_blank"> <b>Slack</b> </a> |
|
||||
</p>
|
||||
|
||||
<div align="center">
|
||||
<img src=assets/perf.png width="90%"/>
|
||||
</div>
|
||||
|
||||
## NEWS
|
||||
- ```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:
|
||||
@@ -91,7 +96,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
|
||||
|
||||
## Distillation and Finetuning
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/training/distillation.html)
|
||||
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetuning.html)
|
||||
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html)
|
||||
|
||||
## 📑 Development Plan
|
||||
|
||||
@@ -111,7 +116,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/developer_guide/overview.html)
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview.html)
|
||||
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects:
|
||||
@@ -128,6 +133,15 @@ We thank MBZUAI and [Anyscale](https://www.anyscale.com/) for their support thro
|
||||
If you use FastVideo for your research, please cite our paper:
|
||||
|
||||
```bibtex
|
||||
@misc{zhang2025vsafastervideodiffusion,
|
||||
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
|
||||
author={Peiyuan Zhang and Haofeng Huang and Yongqi Chen and Will Lin and Zhengzhong Liu and Ion Stoica and Eric Xing and Hao Zhang},
|
||||
year={2025},
|
||||
eprint={2505.13389},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV},
|
||||
url={https://arxiv.org/abs/2505.13389},
|
||||
}
|
||||
@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},
|
||||
|
||||
+42906
-42906
File diff suppressed because it is too large
Load Diff
@@ -1,24 +0,0 @@
|
||||
# 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"
|
||||
@@ -0,0 +1,138 @@
|
||||
# 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/v1/configs/wan_14B_i2v_480p_pipeline.json)
|
||||
- [FastHunyuan-diffusers](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/v1/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.
|
||||
@@ -0,0 +1,5 @@
|
||||
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']
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 1.3 MiB |
Binary file not shown.
|
After Width: | Height: | Size: 8.7 MiB |
Binary file not shown.
|
After Width: | Height: | Size: 769 KiB |
@@ -0,0 +1,645 @@
|
||||
{
|
||||
"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"
|
||||
},
|
||||
"use_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
|
||||
}
|
||||
@@ -0,0 +1,697 @@
|
||||
{
|
||||
"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"
|
||||
},
|
||||
"use_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
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
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, )
|
||||
@@ -0,0 +1,89 @@
|
||||
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, )
|
||||
@@ -0,0 +1,103 @@
|
||||
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
|
||||
@@ -0,0 +1,68 @@
|
||||
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
|
||||
@@ -0,0 +1,24 @@
|
||||
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"
|
||||
}
|
||||
@@ -0,0 +1,38 @@
|
||||
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, )
|
||||
@@ -0,0 +1,88 @@
|
||||
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, )
|
||||
@@ -0,0 +1,315 @@
|
||||
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"
|
||||
}),
|
||||
"use_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,
|
||||
use_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 use_cpu_offload is not None:
|
||||
raw_generation_args['use_cpu_offload'] = use_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")
|
||||
@@ -0,0 +1,593 @@
|
||||
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", "use_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,12 +1,23 @@
|
||||
|
||||
|
||||
# Sliding Tile Atteniton Kernel
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Sliding Tile Attention (STA)
|
||||
We only support H100 for STA.
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup_sta.py install
|
||||
```
|
||||
|
||||
## Installation
|
||||
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
|
||||
First, install C++20 for ThunderKittens:
|
||||
## Video Sparse Attention (VSA)
|
||||
We support H100 (via TK) and RTX 4090 (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
|
||||
@@ -16,17 +27,23 @@ sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
Install STA:
|
||||
(If you use CUDA12.4)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
|
||||
## Usage
|
||||
### STA
|
||||
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.
|
||||
@@ -38,14 +55,22 @@ from st_attn import sliding_tile_attention
|
||||
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)
|
||||
|
||||
```
|
||||
|
||||
### VSA
|
||||
We do not officially supoort end-2-end inference with VSA in FastVideo yet. Stay tuned.
|
||||
|
||||
|
||||
## Test
|
||||
```bash
|
||||
python test/test_sta.py
|
||||
python tests/test_sta.py # test STA
|
||||
python tests/test_block_sparse.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.
|
||||
@@ -5,6 +5,7 @@ import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import torch
|
||||
from st_attn import sliding_tile_attention
|
||||
from triton.testing import do_bench
|
||||
|
||||
|
||||
def flops(batch, seqlen, nheads, headdim, causal, mode="fwd"):
|
||||
@@ -13,16 +14,16 @@ def flops(batch, seqlen, nheads, headdim, causal, mode="fwd"):
|
||||
return f if mode == "fwd" else (2.5 * f if mode == "bwd" else 3.5 * f)
|
||||
|
||||
|
||||
def efficiency(flop, time):
|
||||
flop = flop / 1e12
|
||||
time = time / 1e6
|
||||
return flop / time
|
||||
def compute_TFLOPS(flops, ms):
|
||||
flops = flops / 1e12
|
||||
ms = ms / 1e3
|
||||
return flops / ms
|
||||
|
||||
|
||||
def benchmark_attention(configurations):
|
||||
results = {'fwd': defaultdict(list), 'bwd': defaultdict(list)}
|
||||
|
||||
for B, H, N, D, causal in configurations:
|
||||
for B, H, N, D, causal, dit_seq_shape, window_size in configurations:
|
||||
print("=" * 60)
|
||||
print(f"Timing forward and backward pass for B={B}, H={H}, N={N}, D={D}, causal={causal}")
|
||||
|
||||
@@ -30,38 +31,31 @@ def benchmark_attention(configurations):
|
||||
k = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
|
||||
v = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
|
||||
|
||||
grad_output = torch.randn_like(q, requires_grad=False).contiguous()
|
||||
# grad_output = torch.randn_like(q, requires_grad=False).contiguous()
|
||||
# qg = torch.zeros_like(q, requires_grad=False, dtype=torch.float).contiguous()
|
||||
# kg = torch.zeros_like(k, requires_grad=False, dtype=torch.float).contiguous()
|
||||
# vg = torch.zeros_like(v, requires_grad=False, dtype=torch.float).contiguous()
|
||||
|
||||
qg = torch.zeros_like(q, requires_grad=False, dtype=torch.float).contiguous()
|
||||
kg = torch.zeros_like(k, requires_grad=False, dtype=torch.float).contiguous()
|
||||
vg = torch.zeros_like(v, requires_grad=False, dtype=torch.float).contiguous()
|
||||
|
||||
# Prepare for timing forward pass
|
||||
start_events_fwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
|
||||
end_events_fwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
|
||||
# # Warmup for forward pass
|
||||
# for _ in range(10):
|
||||
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize()
|
||||
# # Time the forward pass
|
||||
# for i in range(10):
|
||||
# start_events_fwd[i].record()
|
||||
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
|
||||
# end_events_fwd[i].record()
|
||||
ms = do_bench(lambda: sliding_tile_attention(q, k, v, [window_size] * 24, 0, False, dit_seq_shape))
|
||||
|
||||
# Warmup for forward pass
|
||||
for _ in range(10):
|
||||
o = sliding_tile_attention(q, k, v, [[6, 6, 6]] * 24, 0, False)
|
||||
# times_fwd = [s.elapsed_time(e) for s, e in zip(start_events_fwd, end_events_fwd)]
|
||||
# time_us_fwd = np.mean(times_fwd) * 1000
|
||||
|
||||
# Time the forward pass
|
||||
for i in range(10):
|
||||
start_events_fwd[i].record()
|
||||
o = sliding_tile_attention(q, k, v, [[6, 6, 6]] * 24, 0, False)
|
||||
end_events_fwd[i].record()
|
||||
|
||||
torch.cuda.synchronize()
|
||||
times_fwd = [s.elapsed_time(e) for s, e in zip(start_events_fwd, end_events_fwd)]
|
||||
time_us_fwd = np.mean(times_fwd) * 1000
|
||||
|
||||
tflops_fwd = efficiency(flops(B, N, H, D, causal, 'fwd'), time_us_fwd)
|
||||
tflops_fwd = compute_TFLOPS(flops(B, N, H, D, causal, 'fwd'), ms)
|
||||
results['fwd'][(D, causal)].append((N, tflops_fwd))
|
||||
|
||||
print(f"Average time for forward pass in us: {time_us_fwd:.2f}")
|
||||
print(f"Average efficiency for forward pass in TFLOPS: {tflops_fwd}")
|
||||
print(f"Average time for forward pass (ms): {ms:.2f}")
|
||||
print(f"Average TFLOPS: {tflops_fwd}")
|
||||
print("-" * 60)
|
||||
|
||||
# torch.cuda.empty_cache()
|
||||
@@ -85,15 +79,14 @@ def benchmark_attention(configurations):
|
||||
# times_bwd = [s.elapsed_time(e) for s, e in zip(start_events_bwd, end_events_bwd)]
|
||||
# time_us_bwd = np.mean(times_bwd) * 1000
|
||||
|
||||
# tflops_bwd = efficiency(flops(B, N, H, D, causal, 'bwd'), time_us_bwd)
|
||||
# tflops_bwd = compute_TFLOPS(flops(B, N, H, D, causal, 'bwd'), ms)
|
||||
# results['bwd'][(D, causal)].append((N, tflops_bwd))
|
||||
|
||||
# print(f"Average time for backward pass in us: {time_us_bwd:.2f}")
|
||||
# print(f"Average efficiency for backward pass in TFLOPS: {tflops_bwd}")
|
||||
print("=" * 60)
|
||||
# print(f"Average time for backward pass(ms): {ms:.2f}")
|
||||
# print(f"Average TFLOPS: {tflops_bwd}")
|
||||
# print("=" * 60)
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
return results
|
||||
|
||||
@@ -124,7 +117,10 @@ def plot_results(results):
|
||||
|
||||
# Example list of configurations to test
|
||||
configurations = [
|
||||
(2, 24, 82944, 128, False),
|
||||
(2, 24, 69120, 128, False, '18x48x80', [3, 6, 10]),
|
||||
(2, 24, 69120, 128, True, '18x48x80', [3, 6, 10]),
|
||||
(2, 24, 82944, 128, False, '36x48x48', [3, 3, 6]), # Stepvideo
|
||||
(2, 24, 82944, 128, True, '36x48x48', [3, 3, 6]),
|
||||
# (16, 16, 768*16, 128, False),
|
||||
# (16, 16, 768*2, 128, False),
|
||||
# (16, 16, 768*4, 128, False),
|
||||
@@ -0,0 +1,225 @@
|
||||
import torch
|
||||
import argparse
|
||||
from flash_attn.utils.benchmark import benchmark_forward
|
||||
from vsa import block_sparse_fwd, block_sparse_bwd
|
||||
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 and backward passes."""
|
||||
print("\n=== BLOCK SPARSE ATTENTION BENCHMARK ===")
|
||||
|
||||
# Forward pass
|
||||
# Warm-up run
|
||||
o, l_vec = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Benchmark forward
|
||||
_, fwd_time = benchmark_forward(
|
||||
block_sparse_fwd,
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num,
|
||||
repeats=20,
|
||||
verbose=False,
|
||||
desc='Block Sparse Forward'
|
||||
)
|
||||
|
||||
sparse_tflops = flops / fwd_time.mean * 1e-12
|
||||
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)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Benchmark backward
|
||||
_, bwd_time = benchmark_forward(
|
||||
block_sparse_bwd,
|
||||
q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num,
|
||||
repeats=20,
|
||||
verbose=False,
|
||||
desc='Block Sparse Backward'
|
||||
)
|
||||
bwd_flops = 2.5 * flops # Approximation
|
||||
|
||||
sparse_bwd_tflops = bwd_flops / bwd_time.mean * 1e-12
|
||||
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()
|
||||
@@ -0,0 +1,217 @@
|
||||
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,6 +1,6 @@
|
||||
### ADD TO THIS TO REGISTER NEW KERNELS
|
||||
sources = {
|
||||
'attn': {
|
||||
'st_attn': {
|
||||
'source_files': {
|
||||
'h100': 'st_attn/st_attn_h100.cu' # define these source files for each GPU target desired.
|
||||
}
|
||||
@@ -9,7 +9,7 @@ sources = {
|
||||
|
||||
### WHICH KERNELS DO WE WANT TO BUILD?
|
||||
# (oftentimes during development work you don't need to redefine them all.)
|
||||
kernels = ['attn']
|
||||
kernels = ['st_attn']
|
||||
|
||||
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
|
||||
target = 'h100'
|
||||
@@ -0,0 +1,15 @@
|
||||
### 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'
|
||||
@@ -0,0 +1,4 @@
|
||||
off_hz = tl.program_id(2)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
from config import kernels, sources, target
|
||||
from csrc.attn.config_sta import kernels, sources, target
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
@@ -0,0 +1,84 @@
|
||||
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.1"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Video Sparse Attention Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/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 = []
|
||||
import torch
|
||||
major, minor = torch.cuda.get_device_capability(0)
|
||||
if major == 9 and minor == 0:# check if H100
|
||||
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"])
|
||||
@@ -7,8 +7,7 @@
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
|
||||
#ifdef TK_COMPILE_ATTN
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
extern torch::Tensor sta_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
|
||||
);
|
||||
@@ -17,8 +16,8 @@ extern torch::Tensor sta_forward(
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "Sliding Block Attention Kernels"; // optional module docstring
|
||||
|
||||
#ifdef TK_COMPILE_ATTN
|
||||
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
|
||||
#endif
|
||||
}
|
||||
}
|
||||
@@ -1,19 +1,22 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
from st_attn_cuda import sta_fwd
|
||||
from torch.utils.checkpoint import detach_variable
|
||||
try:
|
||||
from st_attn_cuda import sta_fwd
|
||||
except ImportError:
|
||||
sta_fwd = None
|
||||
|
||||
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, img_latent_shape='30*48*80'):
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
|
||||
seq_length = q_all.shape[2]
|
||||
img_latent_shape_mapping = {
|
||||
dit_seq_shape_mapping = {
|
||||
'30x48x80':1,
|
||||
'36x48x48':2,
|
||||
'18x48x80':3,
|
||||
}
|
||||
if has_text:
|
||||
assert q_all.shape[
|
||||
2] >= 115200, "STA currently only supports video with latent size (30, 48, 80), which is 117 frames x 768 x 1280 pixels"
|
||||
2] >= 115200 and q_all.shape[2] <= 115456, f"Unsupported {dit_seq_shape}, current shape is {q_all.shape}, only support '30x48x80' for HunyuanVideo"
|
||||
assert q_all.shape[1] == len(window_size), "Number of heads must match the number of window sizes"
|
||||
target_size = math.ceil(seq_length / 384) * 384
|
||||
pad_size = target_size - seq_length
|
||||
@@ -22,14 +25,14 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
|
||||
k_all = torch.cat([k_all, k_all[:, :, -pad_size:]], dim=2)
|
||||
v_all = torch.cat([v_all, v_all[:, :, -pad_size:]], dim=2)
|
||||
else:
|
||||
if img_latent_shape == '36x48x48': # Stepvideo 204x768x68
|
||||
if dit_seq_shape == '36x48x48': # Stepvideo 204x768x68
|
||||
assert q_all.shape[2] == 82944
|
||||
elif img_latent_shape == '18x48x80': # Wan 69x768x1280
|
||||
elif dit_seq_shape == '18x48x80': # Wan 69x768x1280
|
||||
assert q_all.shape[2] == 69120
|
||||
else:
|
||||
raise ValueError(f"Unsupported {img_latent_shape}, current shape is {q_all.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
|
||||
raise ValueError(f"Unsupported {dit_seq_shape}, current shape is {q_all.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
|
||||
|
||||
kernel_aspect_ratio_flag = img_latent_shape_mapping[img_latent_shape]
|
||||
kernel_aspect_ratio_flag = dit_seq_shape_mapping[dit_seq_shape]
|
||||
hidden_states = torch.empty_like(q_all)
|
||||
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
|
||||
for head_index, (t_kernel, h_kernel, w_kernel) in enumerate(window_size):
|
||||
@@ -43,4 +46,4 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
|
||||
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text, kernel_aspect_ratio_flag)
|
||||
if has_text:
|
||||
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True, kernel_aspect_ratio_flag)
|
||||
return hidden_states[:, :, :seq_length]
|
||||
return hidden_states[:, :, :seq_length]
|
||||
+32
-22
@@ -4,9 +4,17 @@
|
||||
#include <cooperative_groups.h>
|
||||
#include <iostream>
|
||||
#include <stdio.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
// #define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
|
||||
__device__ __forceinline__ int clamp_int(int value, int min, int max) {
|
||||
return (value < min) ? min : ((value > max) ? max : value);
|
||||
}
|
||||
// #define ABS(x) ((x) < 0 ? -(x) : (x))
|
||||
__device__ __forceinline__ int abs_int(int value) {
|
||||
return (value < 0) ? -value : value;
|
||||
}
|
||||
|
||||
#define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
|
||||
#define ABS(x) ((x) < 0 ? -(x) : (x))
|
||||
|
||||
constexpr int CONSUMER_WARPGROUPS = (3);
|
||||
constexpr int PRODUCER_WARPGROUPS = (1);
|
||||
@@ -117,16 +125,16 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
|
||||
int qt = seq_idx / 6 / (CH * CW);
|
||||
int qh = (seq_idx / 6) % (CH * CW) / CW;
|
||||
int qw = (seq_idx / 6) % CW;
|
||||
qt = CLAMP(qt, DT, CT-DT-1);
|
||||
qh = CLAMP(qh, DH, CH-DH-1);
|
||||
qw = CLAMP(qw, DW, CW-DW-1);
|
||||
qt = clamp_int(qt, DT, CT-DT-1);
|
||||
qh = clamp_int(qh, DH, CH-DH-1);
|
||||
qw = clamp_int(qw, DW, CW-DW-1);
|
||||
int count = 0;
|
||||
int j = 0;
|
||||
while (count < K::stages - 1) {
|
||||
int kt = j / 3 / (CH * CW);
|
||||
int kh = (j / 3) % (CH * CW) / CW;
|
||||
int kw = (j / 3) % CW;
|
||||
bool mask = (ABS(qt - kt) <= DT) && (ABS(qh - kh) <= DH) && (ABS(qw - kw) <= DW);
|
||||
bool mask = (abs_int(qt - kt) <= DT) && (abs_int(qh - kh) <= DH) && (abs_int(qw - kw) <= DW);
|
||||
if (mask){
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
|
||||
tma::expect_bytes(k_smem_arrived[count], sizeof(k_tile));
|
||||
@@ -167,15 +175,15 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
|
||||
int qt = seq_idx / 6 / (CH * CW);
|
||||
int qh = (seq_idx / 6) % (CH * CW) / CW;
|
||||
int qw = (seq_idx / 6) % CW;
|
||||
qt = CLAMP(qt, DT, CT-DT-1);
|
||||
qh = CLAMP(qh, DH, CH-DH-1);
|
||||
qw = CLAMP(qw, DW, CW-DW-1);
|
||||
int k_t_min = CLAMP(qt-DT, 0, CT-1);
|
||||
int k_t_max = CLAMP(qt+DT, 0, CT-1);
|
||||
int k_h_min = CLAMP(qh-DH, 0, CH-1);
|
||||
int k_h_max = CLAMP(qh+DH, 0, CH-1);
|
||||
int k_w_min = CLAMP(qw-DW, 0, CW-1);
|
||||
int k_w_max = CLAMP(qw+DW, 0, CW-1);
|
||||
qt = clamp_int(qt, DT, CT-DT-1);
|
||||
qh = clamp_int(qh, DH, CH-DH-1);
|
||||
qw = clamp_int(qw, DW, CW-DW-1);
|
||||
int k_t_min = clamp_int(qt-DT, 0, CT-1);
|
||||
int k_t_max = clamp_int(qt+DT, 0, CT-1);
|
||||
int k_h_min = clamp_int(qh-DH, 0, CH-1);
|
||||
int k_h_max = clamp_int(qh+DH, 0, CH-1);
|
||||
int k_w_min = clamp_int(qw-DW, 0, CW-1);
|
||||
int k_w_max = clamp_int(qw+DW, 0, CW-1);
|
||||
int count = 0;
|
||||
for (int kt = k_t_min; kt <= k_t_max; kt++) {
|
||||
for (int kh = k_h_min; kh <= k_h_max; kh++) {
|
||||
@@ -234,7 +242,7 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
|
||||
// the last three kv blocks are for text, we process them separately
|
||||
kv_iters = img_kv_blocks - 1;
|
||||
} else {
|
||||
kv_iters = CLAMP(DT*2+1, 1, CT) * CLAMP(DH*2+1, 1, CH) * CLAMP(DW*2+1, 1, CW) * 3 - 1 ;
|
||||
kv_iters = clamp_int(DT*2+1, 1, CT) * clamp_int(DH*2+1, 1, CH) * clamp_int(DW*2+1, 1, CW) * 3 - 1 ;
|
||||
}
|
||||
|
||||
kittens::wait(qsmem_semaphore, 0);
|
||||
@@ -415,8 +423,9 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
|
||||
float* d_l = reinterpret_cast<float*>(l_ptr);
|
||||
|
||||
cudaDeviceSynchronize();
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
//cudadevicesynchronize();
|
||||
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
|
||||
if (head_dim == 128) {
|
||||
@@ -442,8 +451,8 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
|
||||
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(text_length), static_cast<int>(hr)};
|
||||
|
||||
auto mem_size = kittens::MAX_SHARED_MEMORY;
|
||||
auto threads = NUM_WORKERS * kittens::WARP_THREADS;
|
||||
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
|
||||
int threads = NUM_WORKERS * kittens::WARP_THREADS;
|
||||
if (has_text) {
|
||||
// TORCH_CHECK(seq_len % (CONSUMER_WARPGROUPS*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 192");
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4)-2, qo_heads, batch);
|
||||
@@ -823,9 +832,10 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
|
||||
}
|
||||
CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
cudaStreamSynchronize(stream);
|
||||
// cudaStreamSynchronize(stream);
|
||||
}
|
||||
|
||||
return o;
|
||||
cudaDeviceSynchronize();
|
||||
//cudadevicesynchronize();
|
||||
}
|
||||
|
||||
@@ -0,0 +1,289 @@
|
||||
import torch
|
||||
import argparse
|
||||
from flash_attn.utils.benchmark import benchmark_forward
|
||||
from flash_attn import flash_attn_func
|
||||
from vsa import block_sparse_attn
|
||||
from vsa import BLOCK_M, BLOCK_N
|
||||
|
||||
import numpy as np
|
||||
import random
|
||||
import gc
|
||||
|
||||
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
|
||||
|
||||
|
||||
@torch.no_grad
|
||||
def precision_metric(quant_o, fa2_o):
|
||||
x, xx = quant_o.float(), fa2_o.float()
|
||||
sim = torch.nn.functional.cosine_similarity(x.reshape(1, -1), xx.reshape(1, -1)).item()
|
||||
l1 = ((x - xx).abs().sum() / xx.abs().sum() ).item()
|
||||
rmse = torch.sqrt(torch.mean((x -xx) ** 2)).item()
|
||||
|
||||
return sim, l1, rmse
|
||||
|
||||
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 main(args):
|
||||
set_seed(42)
|
||||
|
||||
# Extract parameters
|
||||
batch = args.batch_size
|
||||
head = args.num_heads
|
||||
headdim = args.head_dim
|
||||
num_iterations = args.num_iterations
|
||||
|
||||
print(f"Block Sparse Attention Benchmark")
|
||||
print(f"batch: {batch}, head: {head}, headdim: {headdim}, iterations: {num_iterations}")
|
||||
|
||||
# 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}")
|
||||
|
||||
# Collect metrics across iterations
|
||||
forward_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_q_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_k_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_v_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
|
||||
for iter_idx in range(num_iterations):
|
||||
if num_iterations > 1:
|
||||
print(f"\nIteration {iter_idx+1}/{num_iterations}")
|
||||
|
||||
# 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)
|
||||
if iter_idx == 0: # Only print this once
|
||||
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, block_sparse_mask = generate_block_sparse_pattern(
|
||||
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
|
||||
|
||||
# expand block_sparse_mask to full mask
|
||||
block_mask_expanded = block_sparse_mask.unsqueeze(-1).unsqueeze(-2) # [b, h, num_q_blocks, num_kv_blocks, 1, 1]
|
||||
block_mask_expanded = block_mask_expanded.expand(-1, -1, -1, -1, BLOCK_M, BLOCK_N) # [b, h, num_q_blocks, num_kv_blocks, BLOCK_M, BLOCK_N]
|
||||
full_mask = block_mask_expanded.permute(0, 1, 2, 4, 3, 5).reshape(batch, head, seq_len, seq_len)
|
||||
|
||||
q.requires_grad = True
|
||||
k.requires_grad = True
|
||||
v.requires_grad = True
|
||||
|
||||
|
||||
# testing forward
|
||||
o = block_sparse_attn(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
del q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask, block_mask_expanded
|
||||
grad_o = torch.randn_like(o)
|
||||
o.backward(grad_o)
|
||||
# clear memory
|
||||
q_sdpa = q.detach().clone()
|
||||
k_sdpa = k.detach().clone()
|
||||
v_sdpa = v.detach().clone()
|
||||
q_sdpa.requires_grad = True
|
||||
k_sdpa.requires_grad = True
|
||||
v_sdpa.requires_grad = True
|
||||
q.data = torch.empty(0, device=q.device)
|
||||
k.data = torch.empty(0, device=k.device)
|
||||
v.data = torch.empty(0, device=v.device)
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
o_sdpa = torch.nn.functional.scaled_dot_product_attention(q_sdpa, k_sdpa, v_sdpa, attn_mask=full_mask)
|
||||
|
||||
|
||||
sim, l1, rmse = precision_metric(o, o_sdpa)
|
||||
assert sim > 0.9999, f"SSIM too low: {sim}"
|
||||
assert l1 < 8e-5, f"l1 too large: {l1}"
|
||||
assert rmse < 2e-5, f"RMSE too large: {rmse}"
|
||||
forward_metrics['sim'].append(sim)
|
||||
forward_metrics['l1'].append(l1)
|
||||
forward_metrics['rmse'].append(rmse)
|
||||
|
||||
print(f"block_sparse_fwd vs torch.nn.functional.scaled_dot_product_attention:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
# test backward
|
||||
o_sdpa.backward(grad_o)
|
||||
|
||||
sim, l1, rmse = precision_metric(q.grad, q_sdpa.grad)
|
||||
# Error bounds collected on H100
|
||||
assert sim > 0.9999, f"SSIM too low: {sim}"
|
||||
assert l1 < 4e-3, f"l1 too large: {l1}"
|
||||
assert rmse < 3e-4, f"RMSE too large: {rmse}"
|
||||
grad_q_metrics['sim'].append(sim)
|
||||
grad_q_metrics['l1'].append(l1)
|
||||
grad_q_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_bwd vs torch.nn.functional.scaled_dot_product_attention grad_q:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
sim, l1, rmse = precision_metric(k.grad, k_sdpa.grad)
|
||||
assert sim > 0.9999, f"SSIM too low: {sim}"
|
||||
assert l1 < 4e-3, f"l1 too large: {l1}"
|
||||
assert rmse < 2e-4, f"RMSE too large: {rmse}"
|
||||
grad_k_metrics['sim'].append(sim)
|
||||
grad_k_metrics['l1'].append(l1)
|
||||
grad_k_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_bwd vs torch.nn.functional.scaled_dot_product_attention grad_k:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
sim, l1, rmse = precision_metric(v.grad, v_sdpa.grad)
|
||||
assert sim > 0.9999, f"SSIM too low: {sim}"
|
||||
assert l1 < 1e-4, f"l1 too large: {l1}"
|
||||
assert rmse < 2e-5, f"RMSE too large: {rmse}"
|
||||
grad_v_metrics['sim'].append(sim)
|
||||
grad_v_metrics['l1'].append(l1)
|
||||
grad_v_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_bwd vs torch.nn.functional.scaled_dot_product_attention grad_v:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
del o, o_sdpa, grad_o, q_sdpa, k_sdpa, v_sdpa
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Print summary statistics if multiple iterations were run
|
||||
if num_iterations > 1:
|
||||
print("\n" + "="*50)
|
||||
print(f"Summary Statistics (over {num_iterations} iterations):")
|
||||
|
||||
print("\nForward metrics:")
|
||||
print(f"Similarity: mean={np.mean(forward_metrics['sim']):.6f}, std={np.std(forward_metrics['sim']):.6f}, min={np.min(forward_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(forward_metrics['l1']):.6f}, std={np.std(forward_metrics['l1']):.6f}, max={np.max(forward_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(forward_metrics['rmse']):.6f}, std={np.std(forward_metrics['rmse']):.6f}, max={np.max(forward_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient Q metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_q_metrics['sim']):.6f}, std={np.std(grad_q_metrics['sim']):.6f}, min={np.min(grad_q_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_q_metrics['l1']):.6f}, std={np.std(grad_q_metrics['l1']):.6f}, max={np.max(grad_q_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_q_metrics['rmse']):.6f}, std={np.std(grad_q_metrics['rmse']):.6f}, max={np.max(grad_q_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient K metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_k_metrics['sim']):.6f}, std={np.std(grad_k_metrics['sim']):.6f}, min={np.min(grad_k_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_k_metrics['l1']):.6f}, std={np.std(grad_k_metrics['l1']):.6f}, max={np.max(grad_k_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_k_metrics['rmse']):.6f}, std={np.std(grad_k_metrics['rmse']):.6f}, max={np.max(grad_k_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient V metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_v_metrics['sim']):.6f}, std={np.std(grad_v_metrics['sim']):.6f}, min={np.min(grad_v_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_v_metrics['l1']):.6f}, std={np.std(grad_v_metrics['l1']):.6f}, max={np.max(grad_v_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_v_metrics['rmse']):.6f}, std={np.std(grad_v_metrics['rmse']):.6f}, max={np.max(grad_v_metrics['rmse']):.6f}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
|
||||
parser.add_argument('--batch_size', type=int, default=4, help='Batch size')
|
||||
parser.add_argument('--num_heads', type=int, default=6, help='Number of heads')
|
||||
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=64, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[29120], help='Sequence lengths to benchmark')
|
||||
parser.add_argument('--num_iterations', type=int, default=50, help='Number of test iterations to run')
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,289 @@
|
||||
import torch
|
||||
import argparse
|
||||
from flash_attn.utils.benchmark import benchmark_forward
|
||||
from flash_attn import flash_attn_func
|
||||
from vsa import triton_attention_sparse
|
||||
from vsa import BLOCK_M, BLOCK_N
|
||||
|
||||
import numpy as np
|
||||
import random
|
||||
import gc
|
||||
|
||||
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
|
||||
|
||||
|
||||
@torch.no_grad
|
||||
def precision_metric(quant_o, fa2_o):
|
||||
x, xx = quant_o.float(), fa2_o.float()
|
||||
sim = torch.nn.functional.cosine_similarity(x.reshape(1, -1), xx.reshape(1, -1)).item()
|
||||
l1 = ((x - xx).abs().sum() / xx.abs().sum() ).item()
|
||||
rmse = torch.sqrt(torch.mean((x -xx) ** 2)).item()
|
||||
|
||||
return sim, l1, rmse
|
||||
|
||||
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 main(args):
|
||||
set_seed(42)
|
||||
|
||||
# Extract parameters
|
||||
batch = args.batch_size
|
||||
head = args.num_heads
|
||||
headdim = args.head_dim
|
||||
num_iterations = args.num_iterations
|
||||
|
||||
print(f"Block Sparse Attention Benchmark")
|
||||
print(f"batch: {batch}, head: {head}, headdim: {headdim}, iterations: {num_iterations}")
|
||||
|
||||
# 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}")
|
||||
|
||||
# Collect metrics across iterations
|
||||
forward_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_q_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_k_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_v_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
|
||||
for iter_idx in range(num_iterations):
|
||||
if num_iterations > 1:
|
||||
print(f"\nIteration {iter_idx+1}/{num_iterations}")
|
||||
|
||||
# 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)
|
||||
if iter_idx == 0: # Only print this once
|
||||
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, block_sparse_mask = generate_block_sparse_pattern(
|
||||
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
|
||||
|
||||
# expand block_sparse_mask to full mask
|
||||
block_mask_expanded = block_sparse_mask.unsqueeze(-1).unsqueeze(-2) # [b, h, num_q_blocks, num_kv_blocks, 1, 1]
|
||||
block_mask_expanded = block_mask_expanded.expand(-1, -1, -1, -1, BLOCK_M, BLOCK_N) # [b, h, num_q_blocks, num_kv_blocks, BLOCK_M, BLOCK_N]
|
||||
full_mask = block_mask_expanded.permute(0, 1, 2, 4, 3, 5).reshape(batch, head, seq_len, seq_len)
|
||||
|
||||
q.requires_grad = True
|
||||
k.requires_grad = True
|
||||
v.requires_grad = True
|
||||
|
||||
|
||||
# testing forward
|
||||
o = triton_attention_sparse(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
del q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask, block_mask_expanded
|
||||
grad_o = torch.randn_like(o)
|
||||
o.backward(grad_o)
|
||||
# clear memory
|
||||
q_sdpa = q.detach().clone()
|
||||
k_sdpa = k.detach().clone()
|
||||
v_sdpa = v.detach().clone()
|
||||
q_sdpa.requires_grad = True
|
||||
k_sdpa.requires_grad = True
|
||||
v_sdpa.requires_grad = True
|
||||
q.data = torch.empty(0, device=q.device)
|
||||
k.data = torch.empty(0, device=k.device)
|
||||
v.data = torch.empty(0, device=v.device)
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
o_sdpa = torch.nn.functional.scaled_dot_product_attention(q_sdpa, k_sdpa, v_sdpa, attn_mask=full_mask)
|
||||
|
||||
|
||||
sim, l1, rmse = precision_metric(o, o_sdpa)
|
||||
assert sim > 0.9999, f"SSIM too low: {sim}"
|
||||
assert l1 < 8e-5, f"l1 too large: {l1}"
|
||||
assert rmse < 5e-5, f"RMSE too large: {rmse}"
|
||||
forward_metrics['sim'].append(sim)
|
||||
forward_metrics['l1'].append(l1)
|
||||
forward_metrics['rmse'].append(rmse)
|
||||
|
||||
print(f"block_sparse_attention_fwd vs torch.nn.functional.scaled_dot_product_attention:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
# test backward
|
||||
o_sdpa.backward(grad_o)
|
||||
|
||||
sim, l1, rmse = precision_metric(q.grad, q_sdpa.grad)
|
||||
# Error bounds collected on H100
|
||||
assert sim > 0.9999, f"SSIM too low: {sim}"
|
||||
assert l1 < 4e-3, f"l1 too large: {l1}"
|
||||
assert rmse < 5e-4, f"RMSE too large: {rmse}"
|
||||
grad_q_metrics['sim'].append(sim)
|
||||
grad_q_metrics['l1'].append(l1)
|
||||
grad_q_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_q:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
sim, l1, rmse = precision_metric(k.grad, k_sdpa.grad)
|
||||
assert sim > 0.9999, f"SSIM too low: {sim}"
|
||||
assert l1 < 4e-3, f"l1 too large: {l1}"
|
||||
assert rmse < 5e-4, f"RMSE too large: {rmse}"
|
||||
grad_k_metrics['sim'].append(sim)
|
||||
grad_k_metrics['l1'].append(l1)
|
||||
grad_k_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_k:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
sim, l1, rmse = precision_metric(v.grad, v_sdpa.grad)
|
||||
assert sim > 0.9999, f"SSIM too low: {sim}"
|
||||
assert l1 < 4e-3, f"l1 too large: {l1}"
|
||||
assert rmse < 5e-4, f"RMSE too large: {rmse}"
|
||||
grad_v_metrics['sim'].append(sim)
|
||||
grad_v_metrics['l1'].append(l1)
|
||||
grad_v_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_v:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
del o, o_sdpa, grad_o, q_sdpa, k_sdpa, v_sdpa
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Print summary statistics if multiple iterations were run
|
||||
if num_iterations > 1:
|
||||
print("\n" + "="*50)
|
||||
print(f"Summary Statistics (over {num_iterations} iterations):")
|
||||
|
||||
print("\nForward metrics:")
|
||||
print(f"Similarity: mean={np.mean(forward_metrics['sim']):.6f}, std={np.std(forward_metrics['sim']):.6f}, min={np.min(forward_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(forward_metrics['l1']):.6f}, std={np.std(forward_metrics['l1']):.6f}, max={np.max(forward_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(forward_metrics['rmse']):.6f}, std={np.std(forward_metrics['rmse']):.6f}, max={np.max(forward_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient Q metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_q_metrics['sim']):.6f}, std={np.std(grad_q_metrics['sim']):.6f}, min={np.min(grad_q_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_q_metrics['l1']):.6f}, std={np.std(grad_q_metrics['l1']):.6f}, max={np.max(grad_q_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_q_metrics['rmse']):.6f}, std={np.std(grad_q_metrics['rmse']):.6f}, max={np.max(grad_q_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient K metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_k_metrics['sim']):.6f}, std={np.std(grad_k_metrics['sim']):.6f}, min={np.min(grad_k_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_k_metrics['l1']):.6f}, std={np.std(grad_k_metrics['l1']):.6f}, max={np.max(grad_k_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_k_metrics['rmse']):.6f}, std={np.std(grad_k_metrics['rmse']):.6f}, max={np.max(grad_k_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient V metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_v_metrics['sim']):.6f}, std={np.std(grad_v_metrics['sim']):.6f}, min={np.min(grad_v_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_v_metrics['l1']):.6f}, std={np.std(grad_v_metrics['l1']):.6f}, max={np.max(grad_v_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_v_metrics['rmse']):.6f}, std={np.std(grad_v_metrics['rmse']):.6f}, max={np.max(grad_v_metrics['rmse']):.6f}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
|
||||
parser.add_argument('--batch_size', type=int, default=4, help='Batch size')
|
||||
parser.add_argument('--num_heads', type=int, default=6, help='Number of heads')
|
||||
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=4, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[4096], help='Sequence lengths to benchmark')
|
||||
parser.add_argument('--num_iterations', type=int, default=10, help='Number of test iterations to run')
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,175 @@
|
||||
import torch
|
||||
from flash_attn_interface import flash_attn_func
|
||||
from st_attn import mha_forward, mha_backward
|
||||
import random
|
||||
from tqdm import tqdm
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
|
||||
def pytorch_test(Q, K, V, dO):
|
||||
q_ = Q.to(torch.float64).requires_grad_()
|
||||
k_ = K.to(torch.float64).requires_grad_()
|
||||
v_ = V.to(torch.float64).requires_grad_()
|
||||
dO_ = dO.to(torch.float64)
|
||||
|
||||
# manual pytorch implementation of scaled dot product attention
|
||||
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_)
|
||||
|
||||
output.backward(dO_)
|
||||
|
||||
q_grad = q_.grad
|
||||
k_grad = k_.grad
|
||||
v_grad = v_.grad
|
||||
|
||||
return output, q_grad, k_grad, v_grad
|
||||
|
||||
def fa2_test(Q, K, V, dO):
|
||||
Q.requires_grad = True
|
||||
K.requires_grad = True
|
||||
V.requires_grad = True
|
||||
output = torch.nn.functional.scaled_dot_product_attention(Q, K, V, is_causal=False)
|
||||
output.backward(dO)
|
||||
|
||||
return output, Q.grad, K.grad, V.grad
|
||||
|
||||
|
||||
def mha_kernel_test(Q, K, V, dO, mode):
|
||||
Q.requires_grad = True
|
||||
K.requires_grad = True
|
||||
V.requires_grad = True
|
||||
|
||||
o, l_vec = mha_forward(Q, K, V)
|
||||
|
||||
if mode == 'forward_only':
|
||||
return o, None, None, None
|
||||
else: # 'forward_backward'
|
||||
qg, kg, vg = mha_backward(Q, K, V, o, l_vec, dO)
|
||||
return o, qg, kg, vg
|
||||
|
||||
def generate_tensor(shape, mean, std, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
|
||||
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
|
||||
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
|
||||
|
||||
return scaled_tensor.contiguous()
|
||||
|
||||
def check_correctness(b, h, n, d, mean, std, num_iterations=100, error_mode='all', test_mode='forward_backward'):
|
||||
results = {
|
||||
'MHA vs PT': {'sum_diff': 0, 'sum_abs': 0, 'max_diff': 0},
|
||||
'FA2 vs PT': {'sum_diff': 0, 'sum_abs': 0, 'max_diff': 0},
|
||||
}
|
||||
|
||||
for _ in range(num_iterations):
|
||||
torch.manual_seed(0)
|
||||
|
||||
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
dO = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
|
||||
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, dO)
|
||||
fa2_o, fa2_qg, fa2_kg, fa2_vg = fa2_test(Q, K, V, dO)
|
||||
|
||||
if test_mode == 'forward_only':
|
||||
mha_o, _, _, _ = mha_kernel_test(Q, K, V, dO, 'forward_only')
|
||||
tensors_mha_pt = [(pt_o, mha_o)]
|
||||
tensors_fa2_pt = [(pt_o, fa2_o)]
|
||||
else: # 'forward_backward'
|
||||
mha_o, mha_qg, mha_kg, mha_vg = mha_kernel_test(Q, K, V, dO, 'forward_backward')
|
||||
|
||||
if error_mode == 'output':
|
||||
tensors_mha_pt = [(pt_o, mha_o)]
|
||||
tensors_fa2_pt = [(pt_o, fa2_o)]
|
||||
elif error_mode == 'backward':
|
||||
tensors_mha_pt = [(pt_qg, mha_qg),
|
||||
(pt_kg, mha_kg),
|
||||
(pt_vg, mha_vg)]
|
||||
tensors_fa2_pt = [(pt_qg, fa2_qg),
|
||||
(pt_kg, fa2_kg),
|
||||
(pt_vg, fa2_vg)]
|
||||
else: # 'all'
|
||||
tensors_mha_pt = [(pt_o, mha_o),
|
||||
(pt_qg, mha_qg),
|
||||
(pt_kg, mha_kg),
|
||||
(pt_vg, mha_vg)]
|
||||
tensors_fa2_pt = [(pt_o, fa2_o),
|
||||
(pt_qg, fa2_qg),
|
||||
(pt_kg, fa2_kg),
|
||||
(pt_vg, fa2_vg)]
|
||||
|
||||
for pt, mha in tensors_mha_pt:
|
||||
diff = pt - mha
|
||||
abs_diff = torch.abs(diff)
|
||||
results['MHA vs PT']['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results['MHA vs PT']['sum_abs'] += torch.sum(torch.abs(pt)).item()
|
||||
results['MHA vs PT']['max_diff'] = max(results['MHA vs PT']['max_diff'], torch.max(abs_diff).item())
|
||||
|
||||
for pt, fa2 in tensors_fa2_pt:
|
||||
diff = pt - fa2
|
||||
abs_diff = torch.abs(diff)
|
||||
results['FA2 vs PT']['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results['FA2 vs PT']['sum_abs'] += torch.sum(torch.abs(pt)).item()
|
||||
results['FA2 vs PT']['max_diff'] = max(results['FA2 vs PT']['max_diff'], torch.max(abs_diff).item())
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Calculate total elements based on test mode and error mode
|
||||
if test_mode == 'forward_only':
|
||||
total_elements = b * h * n * d * num_iterations
|
||||
else: # 'forward_backward'
|
||||
total_elements = b * h * n * d * num_iterations * (1 if error_mode == 'output' else 3 if error_mode == 'backward' else 4)
|
||||
|
||||
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_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward'):
|
||||
seq_lengths = [768 * (2**i) for i in range(1)]
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print(f"MHA ERROR COMPARISON TABLE (b={b}, h={h}, d={d}, mean={mean}, std={std})")
|
||||
print(f"Mode: {error_mode}, Test: {test_mode}")
|
||||
print(f"{'='*80}")
|
||||
|
||||
# Print header
|
||||
print(f"{'Seq Length':<12} | {'MHA vs PT Avg':<15} | {'MHA vs PT Max':<15} | {'FA2 vs PT Avg':<15} | {'FA2 vs PT Max':<15}")
|
||||
print(f"{'-'*12} | {'-'*15} | {'-'*15} | {'-'*15} | {'-'*15}")
|
||||
|
||||
for n in seq_lengths:
|
||||
results = check_correctness(b, h, n, d, mean, std, error_mode=error_mode, test_mode=test_mode)
|
||||
|
||||
mha_pt_avg = results['MHA vs PT']['avg_diff']
|
||||
mha_pt_max = results['MHA vs PT']['max_diff']
|
||||
fa2_pt_avg = results['FA2 vs PT']['avg_diff']
|
||||
fa2_pt_max = results['FA2 vs PT']['max_diff']
|
||||
|
||||
# Print row
|
||||
print(f"{n:<12} | {mha_pt_avg:<15.6e} | {mha_pt_max:<15.6e} | {fa2_pt_avg:<15.6e} | {fa2_pt_max:<15.6e}")
|
||||
|
||||
print(f"{'='*80}\n")
|
||||
|
||||
# fix random seed
|
||||
torch.manual_seed(0)
|
||||
|
||||
# Example usage
|
||||
b, h, d = 2, 2, 64
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
# Test forward only
|
||||
generate_error_tables(b, h, d, mean, std, error_mode='output', test_mode='forward_only')
|
||||
|
||||
# Test forward and backward
|
||||
generate_error_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward')
|
||||
|
||||
print("MHA attention error comparison completed.")
|
||||
@@ -2,27 +2,28 @@ import torch
|
||||
from flex_sta_ref import get_sliding_tile_attention_mask
|
||||
from st_attn import sliding_tile_attention
|
||||
from torch.nn.attention.flex_attention import flex_attention
|
||||
# from flash_attn_interface import flash_attn_func
|
||||
from tqdm import tqdm
|
||||
|
||||
flex_attention = torch.compile(flex_attention, dynamic=False)
|
||||
|
||||
|
||||
def flex_test(Q, K, V, kernel_size):
|
||||
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (36, 48, 48), 39, 'cuda', 0)
|
||||
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (18, 48, 80), 0, 'cuda', 0)
|
||||
output = flex_attention(Q, K, V, block_mask=mask)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def h100_fwd_kernel_test(Q, K, V, kernel_size):
|
||||
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 39, False)
|
||||
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
|
||||
return o
|
||||
|
||||
|
||||
def generate_tensor(shape, mean, std, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
|
||||
magnitude = torch.linalg.norm(tensor, dim=-1, keepdim=True)
|
||||
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
|
||||
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
|
||||
|
||||
return scaled_tensor.contiguous()
|
||||
@@ -36,7 +37,7 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
|
||||
'max_diff': 0
|
||||
},
|
||||
}
|
||||
kernel_size_ls = [(6, 1, 6), (6, 6, 1)]
|
||||
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
|
||||
from tqdm import tqdm
|
||||
for kernel_size in tqdm(kernel_size_ls):
|
||||
for _ in range(num_iterations):
|
||||
@@ -71,25 +72,16 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
|
||||
return results
|
||||
|
||||
|
||||
def generate_error_graphs(b, h, d, causal, mean, std, error_mode='all'):
|
||||
seq_lengths = [82944]
|
||||
|
||||
tk_avg_errors, tk_max_errors = [], []
|
||||
|
||||
for n in tqdm(seq_lengths, desc="Generating error data"):
|
||||
results = check_correctness(b, h, n, d, causal, mean, std, error_mode=error_mode)
|
||||
|
||||
tk_avg_errors.append(results['TK vs FLEX']['avg_diff'])
|
||||
tk_max_errors.append(results['TK vs FLEX']['max_diff'])
|
||||
|
||||
|
||||
# Example usage
|
||||
b, h, d = 2, 24, 128
|
||||
n = 69120 # Sequence length
|
||||
causal = False
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
for mode in ['output']:
|
||||
generate_error_graphs(b, h, d, causal, mean, std, error_mode=mode)
|
||||
|
||||
print("Error graphs generated and saved for all modes.")
|
||||
# Run correctness check directly
|
||||
results = check_correctness(b, h, n, d, causal, mean, std, error_mode='output')
|
||||
assert results['TK vs FLEX']['avg_diff'] < 3e-6, f"Average difference: {results['TK vs FLEX']['avg_diff']} is too large"
|
||||
assert results['TK vs FLEX']['max_diff'] < 4e-2, f"Maximum difference: {results['TK vs FLEX']['max_diff']} is too large"
|
||||
print(f"Average difference: {results['TK vs FLEX']['avg_diff']}")
|
||||
print(f"Maximum difference: {results['TK vs FLEX']['max_diff']}")
|
||||
@@ -0,0 +1,183 @@
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
|
||||
from vsa import triton_attention
|
||||
|
||||
def pytorch_test(Q, K, V, dO):
|
||||
q_ = Q.to(torch.float64).requires_grad_()
|
||||
k_ = K.to(torch.float64).requires_grad_()
|
||||
v_ = V.to(torch.float64).requires_grad_()
|
||||
q_.grad = None
|
||||
k_.grad = None
|
||||
v_.grad = None
|
||||
dO_ = dO.to(torch.float64)
|
||||
|
||||
# manual pytorch implementation of scaled dot product attention
|
||||
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_)
|
||||
|
||||
output.backward(dO_)
|
||||
|
||||
q_grad = q_.grad
|
||||
k_grad = k_.grad
|
||||
v_grad = v_.grad
|
||||
|
||||
return output, q_grad, k_grad, v_grad
|
||||
|
||||
def fa2_test(Q, K, V, dO):
|
||||
Q.requires_grad = True
|
||||
K.requires_grad = True
|
||||
V.requires_grad = True
|
||||
Q.grad = None
|
||||
K.grad = None
|
||||
V.grad = None
|
||||
|
||||
output = torch.nn.functional.scaled_dot_product_attention(Q, K, V, is_causal=False)
|
||||
output.backward(dO)
|
||||
|
||||
return output, Q.grad, K.grad, V.grad
|
||||
|
||||
def triton_test(Q, K, V, dO):
|
||||
Q.requires_grad = True
|
||||
K.requires_grad = True
|
||||
V.requires_grad = True
|
||||
Q.grad = None
|
||||
K.grad = None
|
||||
V.grad = None
|
||||
|
||||
output = triton_attention(Q, K, V)
|
||||
|
||||
output.backward(dO)
|
||||
|
||||
q_grad = Q.grad
|
||||
k_grad = K.grad
|
||||
v_grad = V.grad
|
||||
|
||||
return output.to(Q.dtype) if output is not None else None, q_grad, k_grad, v_grad
|
||||
|
||||
def generate_tensor(shape, mean, std, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
|
||||
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
|
||||
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
|
||||
|
||||
return scaled_tensor.contiguous()
|
||||
|
||||
def check_correctness(b, h, n, d, mean, std, num_iterations=100, error_mode='all', test_mode='forward_backward'):
|
||||
results = {
|
||||
'FA2 vs PT': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
'Triton vs PT': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
}
|
||||
|
||||
for _ in range(num_iterations):
|
||||
torch.manual_seed(0)
|
||||
|
||||
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
dO = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
|
||||
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, dO)
|
||||
fa2_o, fa2_qg, fa2_kg, fa2_vg = fa2_test(Q, K, V, dO)
|
||||
triton_o, triton_qg, triton_kg, triton_vg = triton_test(Q, K, V, dO)
|
||||
|
||||
if test_mode == 'forward_only':
|
||||
tensors_fa2_pt = [(pt_o, fa2_o)]
|
||||
tensors_triton_pt = [(pt_o, triton_o)]
|
||||
else: # 'forward_backward'
|
||||
if error_mode == 'output':
|
||||
tensors_fa2_pt = [(pt_o, fa2_o)]
|
||||
tensors_triton_pt = [(pt_o, triton_o)]
|
||||
elif error_mode == 'backward':
|
||||
tensors_fa2_pt = [(pt_qg, fa2_qg),
|
||||
(pt_kg, fa2_kg),
|
||||
(pt_vg, fa2_vg)]
|
||||
tensors_triton_pt = [(pt_qg, triton_qg),
|
||||
(pt_kg, triton_kg),
|
||||
(pt_vg, triton_vg)]
|
||||
else: # 'all'
|
||||
tensors_fa2_pt = [(pt_o, fa2_o),
|
||||
(pt_qg, fa2_qg),
|
||||
(pt_kg, fa2_kg),
|
||||
(pt_vg, fa2_vg)]
|
||||
tensors_triton_pt = [(pt_o, triton_o),
|
||||
(pt_qg, triton_qg),
|
||||
(pt_kg, triton_kg),
|
||||
(pt_vg, triton_vg)]
|
||||
|
||||
for pt, fa2 in tensors_fa2_pt:
|
||||
diff = pt - fa2
|
||||
abs_diff = torch.abs(diff)
|
||||
results['FA2 vs PT']['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results['FA2 vs PT']['sum_abs'] += torch.sum(torch.abs(pt)).item()
|
||||
results['FA2 vs PT']['max_diff'] = max(results['FA2 vs PT']['max_diff'], torch.max(abs_diff).item())
|
||||
|
||||
for pt, triton in tensors_triton_pt:
|
||||
diff = pt - triton
|
||||
abs_diff = torch.abs(diff)
|
||||
results['Triton vs PT']['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results['Triton vs PT']['sum_abs'] += torch.sum(torch.abs(pt)).item()
|
||||
results['Triton vs PT']['max_diff'] = max(results['Triton vs PT']['max_diff'], torch.max(abs_diff).item())
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Calculate total elements based on test mode and error mode
|
||||
if test_mode == 'forward_only':
|
||||
total_elements = b * h * n * d * num_iterations
|
||||
else: # 'forward_backward'
|
||||
total_elements = b * h * n * d * num_iterations * (1 if error_mode == 'output' else 3 if error_mode == 'backward' else 4)
|
||||
|
||||
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_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward'):
|
||||
seq_lengths = [768 * (2**i) for i in range(1)]
|
||||
|
||||
print(f"\n{'='*100}")
|
||||
print(f"ATTENTION ERROR COMPARISON TABLE (b={b}, h={h}, d={d}, mean={mean}, std={std})")
|
||||
print(f"Mode: {error_mode}, Test: {test_mode}")
|
||||
print(f"{'='*100}")
|
||||
|
||||
# Print header
|
||||
print(f"{'Seq Length':<12} | {'FA2 vs PT Avg':<15} | {'FA2 vs PT Max':<15} | {'Triton vs PT Avg':<15} | {'Triton vs PT Max':<15}")
|
||||
print(f"{'-'*12} | {'-'*15} | {'-'*15} | {'-'*15} | {'-'*15}")
|
||||
|
||||
for n in seq_lengths:
|
||||
results = check_correctness(b, h, n, d, mean, std, error_mode=error_mode, test_mode=test_mode)
|
||||
|
||||
fa2_pt_avg = results['FA2 vs PT']['avg_diff']
|
||||
fa2_pt_max = results['FA2 vs PT']['max_diff']
|
||||
triton_pt_avg = results['Triton vs PT']['avg_diff']
|
||||
triton_pt_max = results['Triton vs PT']['max_diff']
|
||||
|
||||
# Print row with both comparisons
|
||||
print(f"{n:<12} | {fa2_pt_avg:<15.6e} | {fa2_pt_max:<15.6e} | {triton_pt_avg:<15.6e} | {triton_pt_max:<15.6e}")
|
||||
|
||||
print(f"{'='*100}\n")
|
||||
|
||||
# fix random seed
|
||||
torch.manual_seed(0)
|
||||
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
configs = [
|
||||
(4, 1, 128), # Larger batch, single head, larger dim
|
||||
(2, 8, 64), # Medium batch, many heads, medium dim
|
||||
]
|
||||
|
||||
for b, h, d in configs:
|
||||
print(f"\nConfiguration: batch={b}, heads={h}, dim={d}")
|
||||
generate_error_tables(b, h, d, mean, std, error_mode='backward', test_mode='forward_backward')
|
||||
|
||||
print("Attention error comparison completed.")
|
||||
@@ -0,0 +1,27 @@
|
||||
#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
|
||||
);
|
||||
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
|
||||
);
|
||||
#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
|
||||
}
|
||||
@@ -0,0 +1,278 @@
|
||||
import torch
|
||||
from typing import Tuple
|
||||
from vsa.vsa import block_sparse_attn
|
||||
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 attention as triton_attention, attention_sparse as triton_attention_sparse
|
||||
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, 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]
|
||||
|
||||
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 == 0 and 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).mean(dim=3)
|
||||
k_compress = k.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).mean(dim=3)
|
||||
v_compress = v.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).mean(dim=3)
|
||||
|
||||
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)
|
||||
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num = generate_topk_block_sparse_pattern(
|
||||
block_attn_score, topk)
|
||||
|
||||
output_select = block_sparse_attn(q, k, v, q2k_block_sparse_index,
|
||||
q2k_block_sparse_num,
|
||||
k2q_block_sparse_index,
|
||||
k2q_block_sparse_num)
|
||||
|
||||
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
|
||||
|
||||
|
||||
def generate_topk_block_sparse_pattern(block_attn_score: torch.Tensor,
|
||||
topk: int):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly topk kv blocks,
|
||||
based on the provided attention scores.
|
||||
|
||||
Args:
|
||||
block_attn_score: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Attention scores between query and key blocks
|
||||
topk: int
|
||||
Number of kv blocks each q block attends to
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, topk]
|
||||
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 topk).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, max_q_per_kv]
|
||||
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.
|
||||
"""
|
||||
device = block_attn_score.device
|
||||
# Extract dimensions from block_attn_score
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_attn_score.shape
|
||||
|
||||
sorted_result = torch.sort(block_attn_score, dim=-1, descending=True)
|
||||
|
||||
sorted_indice = sorted_result.indices
|
||||
|
||||
q2k_block_sparse_index, _ = torch.sort(sorted_indice[:, :, :, :topk],
|
||||
dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(dtype=torch.int32)
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks),
|
||||
topk,
|
||||
device=device,
|
||||
dtype=torch.int32)
|
||||
|
||||
block_map = topk_index_to_map(q2k_block_sparse_index,
|
||||
num_kv_blocks,
|
||||
transpose_map=True)
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(
|
||||
block_map.transpose(2, 3))
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
|
||||
|
||||
## pytorch sdpa version of block sparse ##
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
|
||||
@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: tl.constexpr,
|
||||
):
|
||||
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: tl.constexpr,
|
||||
):
|
||||
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.static_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
|
||||
|
||||
@@ -0,0 +1,707 @@
|
||||
"""
|
||||
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 ─────────────────────────────
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_fwd_inner(acc, l_i, m_i, q, #
|
||||
K_block_ptr, V_block_ptr, #
|
||||
start_m, qk_scale, #
|
||||
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr, BLOCK_N: tl.constexpr, #
|
||||
STAGE: tl.constexpr, offs_m: tl.constexpr, offs_n: tl.constexpr, #
|
||||
N_CTX: tl.constexpr, fp8_v: tl.constexpr):
|
||||
|
||||
# loop over k, v and update accumulator
|
||||
for start_n in range(0, N_CTX, BLOCK_N):
|
||||
# -- compute qk ----
|
||||
k = tl.load(K_block_ptr)
|
||||
qk = tl.dot(q, k)
|
||||
|
||||
m_ij = tl.maximum(m_i, tl.max(qk, 1) * qk_scale)
|
||||
qk = qk * qk_scale - m_ij[:, None]
|
||||
p = tl.math.exp2(qk)
|
||||
l_ij = tl.sum(p, 1)
|
||||
# -- update m_i and l_i
|
||||
alpha = tl.math.exp2(m_i - m_ij)
|
||||
l_i = l_i * alpha + l_ij
|
||||
# -- update output accumulator --
|
||||
acc = acc * alpha[:, None]
|
||||
# update acc
|
||||
v = tl.load(V_block_ptr)
|
||||
if fp8_v:
|
||||
p = p.to(tl.float8e5)
|
||||
else:
|
||||
p = p.to(tl.bfloat16)
|
||||
acc = tl.dot(p, v, acc)
|
||||
# update m_i and l_i
|
||||
m_i = m_ij
|
||||
V_block_ptr = tl.advance(V_block_ptr, (BLOCK_N, 0))
|
||||
K_block_ptr = tl.advance(K_block_ptr, (0, BLOCK_N))
|
||||
return acc, l_i, m_i
|
||||
|
||||
|
||||
# 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, #
|
||||
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)
|
||||
|
||||
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)
|
||||
|
||||
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.autotune(configs, key=["N_CTX", "HEAD_DIM"])
|
||||
@triton.jit
|
||||
def _attn_fwd(Q, K, V, sm_scale, 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 #
|
||||
):
|
||||
tl.static_assert(BLOCK_N <= HEAD_DIM)
|
||||
start_m = tl.program_id(0)
|
||||
off_hz = tl.program_id(1)
|
||||
off_z = off_hz // H
|
||||
off_h = off_hz % H
|
||||
qvk_offset = off_z.to(tl.int64) * stride_qz + off_h.to(tl.int64) * stride_qh
|
||||
|
||||
# block pointers
|
||||
Q_block_ptr = tl.make_block_ptr(
|
||||
base=Q + qvk_offset,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_qm, stride_qk),
|
||||
offsets=(start_m * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1, 0)
|
||||
V_block_ptr = tl.make_block_ptr(
|
||||
base=V + qvk_offset,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_vk, stride_vn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(BLOCK_N, HEAD_DIM),
|
||||
order=v_order,
|
||||
)
|
||||
K_block_ptr = tl.make_block_ptr(
|
||||
base=K + qvk_offset,
|
||||
shape=(HEAD_DIM, N_CTX),
|
||||
strides=(stride_kk, stride_kn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_N),
|
||||
order=(0, 1),
|
||||
)
|
||||
O_block_ptr = tl.make_block_ptr(
|
||||
base=Out + qvk_offset,
|
||||
shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_om, stride_on),
|
||||
offsets=(start_m * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM),
|
||||
order=(1, 0),
|
||||
)
|
||||
# initialize offsets
|
||||
offs_m = start_m * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
offs_n = tl.arange(0, BLOCK_N)
|
||||
# initialize pointer to m and l
|
||||
m_i = tl.zeros([BLOCK_M], dtype=tl.float32) - float("inf")
|
||||
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
|
||||
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
|
||||
# load scales
|
||||
qk_scale = sm_scale
|
||||
qk_scale *= 1.44269504 # 1/log(2)
|
||||
# load q: it will stay in SRAM throughout
|
||||
q = tl.load(Q_block_ptr)
|
||||
acc, l_i, m_i = _attn_fwd_inner(acc, l_i, m_i, q, K_block_ptr, V_block_ptr, #
|
||||
start_m, qk_scale, #
|
||||
BLOCK_M, HEAD_DIM, BLOCK_N, #
|
||||
3, offs_m, offs_n, N_CTX, V.dtype.element_ty == tl.float8e5 #
|
||||
)
|
||||
|
||||
# epilogue
|
||||
m_i += tl.math.log2(l_i)
|
||||
acc = acc / l_i[:, None]
|
||||
m_ptrs = M + off_hz * N_CTX + offs_m
|
||||
tl.store(m_ptrs, m_i)
|
||||
tl.store(O_block_ptr, acc.to(Out.type.element_ty))
|
||||
|
||||
|
||||
@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,
|
||||
# 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
|
||||
|
||||
|
||||
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, :])
|
||||
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,
|
||||
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
|
||||
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)
|
||||
# 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,
|
||||
# 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,
|
||||
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,
|
||||
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)
|
||||
|
||||
|
||||
|
||||
class _attention(torch.autograd.Function):
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, q, k, v):
|
||||
# shape constraints
|
||||
sm_scale = 1.0 / (q.shape[-1] ** 0.5)
|
||||
HEAD_DIM_Q, HEAD_DIM_K = q.shape[-1], k.shape[-1]
|
||||
# when v is in float8_e5m2 it is transposed.
|
||||
HEAD_DIM_V = v.shape[-1]
|
||||
assert HEAD_DIM_Q == HEAD_DIM_K and HEAD_DIM_K == HEAD_DIM_V
|
||||
assert HEAD_DIM_K in {16, 32, 64, 128, 256}
|
||||
o = torch.empty_like(q)
|
||||
stage = 1
|
||||
extra_kern_args = {}
|
||||
|
||||
|
||||
grid = lambda args: (triton.cdiv(q.shape[2], args["BLOCK_M"]), q.shape[0] * q.shape[1], 1)
|
||||
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
_attn_fwd[grid](
|
||||
q, k, v, sm_scale, 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), #
|
||||
q.shape[0], q.shape[1], #
|
||||
N_CTX=q.shape[2], #
|
||||
HEAD_DIM=HEAD_DIM_K, #
|
||||
STAGE=stage, #
|
||||
**extra_kern_args)
|
||||
|
||||
ctx.save_for_backward(q, k, v, o, M)
|
||||
ctx.grid = grid
|
||||
ctx.sm_scale = sm_scale
|
||||
ctx.HEAD_DIM = HEAD_DIM_K
|
||||
return o
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, do):
|
||||
q, k, v, o, M = ctx.saved_tensors
|
||||
assert do.is_contiguous()
|
||||
assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
BATCH, N_HEAD, N_CTX = q.shape[:3]
|
||||
PRE_BLOCK = 128
|
||||
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 * (ctx.sm_scale * RCP_LN2)
|
||||
PRE_BLOCK = 128
|
||||
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=ctx.HEAD_DIM #
|
||||
)
|
||||
grid = (N_CTX // BLOCK_N1, 1, BATCH * N_HEAD)
|
||||
_attn_bwd[grid](
|
||||
q, arg_k, v, ctx.sm_scale, do, dq, dk, dv, #
|
||||
M, delta, #
|
||||
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=ctx.HEAD_DIM #
|
||||
)
|
||||
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
attention = _attention.apply
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
class _attention_sparse(torch.autograd.Function):
|
||||
"""
|
||||
Thin autograd wrapper that uses the sparse forward kernel above and the
|
||||
standard dense backward kernels defined earlier (no extra memory use).
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def forward(ctx, q, k, v, q2k_index, q2k_num, k2q_index, k2q_num):
|
||||
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,
|
||||
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
|
||||
)
|
||||
|
||||
ctx.save_for_backward(q, k, v, o, M, q2k_index, q2k_num, k2q_index, k2q_num)
|
||||
ctx.grid = None
|
||||
ctx.sm_scale = sm_scale
|
||||
ctx.HEAD_DIM = D
|
||||
return o
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, do):
|
||||
q, k, v, o, M, q2k_index, q2k_num, k2q_index, k2q_num = ctx.saved_tensors
|
||||
assert do.is_contiguous()
|
||||
assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
BATCH, N_HEAD, N_CTX = q.shape[:3]
|
||||
PRE_BLOCK = 128
|
||||
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 * (ctx.sm_scale * RCP_LN2)
|
||||
PRE_BLOCK = 128
|
||||
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=ctx.HEAD_DIM #
|
||||
)
|
||||
|
||||
|
||||
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, ctx.sm_scale, do, dq, dk, dv, #
|
||||
M, delta, #
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
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=ctx.HEAD_DIM #
|
||||
)
|
||||
|
||||
return dq, dk, dv, None, None, None, None
|
||||
|
||||
attention_sparse = _attention_sparse.apply
|
||||
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
|
||||
|
||||
|
||||
try:
|
||||
from flash_attn.flash_attn_interface import \
|
||||
flash_attn_qkvpacked_func as flash_attn_func
|
||||
HAS_FLASH = True
|
||||
except BaseException:
|
||||
HAS_FLASH = False
|
||||
|
||||
TORCH_HAS_FP8 = hasattr(torch, 'float8_e5m2')
|
||||
BATCH, N_HEADS, HEAD_DIM = 4, 32, 128
|
||||
# vary seq length for fixed head and batch=4
|
||||
configs = []
|
||||
for mode in ["fwd", "bwd"]:
|
||||
configs.append(
|
||||
triton.testing.Benchmark(
|
||||
x_names=["N_CTX"],
|
||||
x_vals=[2**i for i in range(10, 15)],
|
||||
line_arg="provider",
|
||||
line_vals=["triton-fp16"] + (["triton-fp8"] if TORCH_HAS_FP8 else []) +
|
||||
(["flash"] if HAS_FLASH else []),
|
||||
line_names=["Triton [FP16]"] + (["Triton [FP8]"] if TORCH_HAS_FP8 else []) +
|
||||
(["Flash-2"] if HAS_FLASH else []),
|
||||
styles=[("red", "-"), ("blue", "-"), ("green", "-")],
|
||||
ylabel="ms",
|
||||
plot_name=f"fused-attention-batch{BATCH}-head{N_HEADS}-d{HEAD_DIM}-{mode}",
|
||||
args={
|
||||
"H": N_HEADS,
|
||||
"BATCH": BATCH,
|
||||
"HEAD_DIM": HEAD_DIM,
|
||||
"mode": mode,
|
||||
},
|
||||
))
|
||||
|
||||
|
||||
@triton.testing.perf_report(configs)
|
||||
def bench_flash_attention(BATCH, H, N_CTX, HEAD_DIM, mode, provider, device="cuda"):
|
||||
assert mode in ["fwd", "bwd"]
|
||||
warmup = 25
|
||||
rep = 100
|
||||
dtype = torch.bfloat16
|
||||
if "triton" in provider:
|
||||
q = torch.randn((BATCH, H, N_CTX, HEAD_DIM), dtype=dtype, device=device, requires_grad=True)
|
||||
k = torch.randn((BATCH, H, N_CTX, HEAD_DIM), dtype=dtype, device=device, requires_grad=True)
|
||||
v = torch.randn((BATCH, H, N_CTX, HEAD_DIM), dtype=dtype, device=device, requires_grad=True)
|
||||
if mode == "fwd" and "fp8" in provider:
|
||||
q = q.to(torch.float8_e5m2)
|
||||
k = k.to(torch.float8_e5m2)
|
||||
v = v.permute(0, 1, 3, 2).contiguous()
|
||||
v = v.permute(0, 1, 3, 2)
|
||||
v = v.to(torch.float8_e5m2)
|
||||
fn = lambda: attention(q, k, v)
|
||||
if mode == "bwd":
|
||||
o = fn()
|
||||
do = torch.randn_like(o)
|
||||
fn = lambda: o.backward(do, retain_graph=True)
|
||||
ms = triton.testing.do_bench(fn, warmup=warmup, rep=rep)
|
||||
if provider == "flash":
|
||||
qkv = torch.randn((BATCH, N_CTX, 3, H, HEAD_DIM), dtype=dtype, device=device, requires_grad=True)
|
||||
fn = lambda: flash_attn_func(qkv, causal=False)
|
||||
if mode == "bwd":
|
||||
o = fn()
|
||||
do = torch.randn_like(o)
|
||||
fn = lambda: o.backward(do, retain_graph=True)
|
||||
ms = triton.testing.do_bench(fn, warmup=warmup, rep=rep)
|
||||
flops_per_matmul = 2.0 * BATCH * H * N_CTX * N_CTX * HEAD_DIM
|
||||
total_flops = 2 * flops_per_matmul
|
||||
if mode == "bwd":
|
||||
total_flops *= 2.5 # 2.0(bwd) + 0.5(recompute)
|
||||
return total_flops / ms * 1e-9
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
# only works on post-Ampere GPUs right now
|
||||
bench_flash_attention.run(save_path=".", print_data=True)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,47 @@
|
||||
import torch
|
||||
try:
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
except ImportError:
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
from .block_sparse_attn_triton import attention_sparse as block_sparse_attn_triton
|
||||
|
||||
class BlockSparseAttentionFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
o, lse = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
ctx.save_for_backward(q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
return o
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num = ctx.saved_tensors
|
||||
grad_q, grad_k, grad_v = block_sparse_bwd(
|
||||
q, k, v, o, lse, grad_output, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
)
|
||||
return grad_q, grad_k, grad_v, None, None, None, None
|
||||
|
||||
|
||||
@torch._dynamo.disable
|
||||
def block_sparse_attn(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
"""
|
||||
Differentiable block sparse attention function.
|
||||
|
||||
Args:
|
||||
q: Query tensor [batch_size, num_heads, seq_len_q, head_dim]
|
||||
k: Key tensor [batch_size, num_heads, seq_len_kv, head_dim]
|
||||
v: Value tensor [batch_size, num_heads, seq_len_kv, head_dim]
|
||||
q2k_block_sparse_index: Indices for query-to-key sparse blocks
|
||||
q2k_block_sparse_num: Number of sparse blocks for each query block
|
||||
k2q_block_sparse_index: Indices for key-to-query sparse blocks (for backward pass)
|
||||
k2q_block_sparse_num: Number of sparse blocks for each key block (for backward pass)
|
||||
|
||||
Returns:
|
||||
output: Attention output tensor [batch_size, num_heads, seq_len_q, head_dim]
|
||||
"""
|
||||
if block_sparse_fwd is not None:
|
||||
return BlockSparseAttentionFunction.apply(
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
)
|
||||
else:
|
||||
return block_sparse_attn_triton(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
@@ -1,7 +1,9 @@
|
||||
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
|
||||
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 \
|
||||
@@ -9,17 +11,25 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
zsh \
|
||||
vim \
|
||||
curl \
|
||||
gcc-11 \
|
||||
g++-11 \
|
||||
clang-11 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
|
||||
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
|
||||
rm Miniconda3-latest-Linux-x86_64.sh
|
||||
# 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
|
||||
|
||||
ENV PATH=/opt/conda/bin:$PATH
|
||||
# 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
|
||||
|
||||
RUN conda create --name fastvideo-dev python=3.12.9 -y
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
# 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 ./
|
||||
@@ -27,22 +37,36 @@ COPY pyproject.toml ./
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
|
||||
conda clean -afy
|
||||
# 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.0.post2 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
|
||||
# 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
|
||||
|
||||
# Remove authentication headers
|
||||
RUN git config --unset-all http.https://github.com/.extraheader || true
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_sta.py install
|
||||
|
||||
# Set up automatic conda environment activation for all shells
|
||||
RUN echo 'source /opt/conda/etc/profile.d/conda.sh' >> /root/.bashrc && \
|
||||
echo 'conda activate fastvideo-dev' >> /root/.bashrc && \
|
||||
# Ensure .bashrc is sourced for SSH login shells
|
||||
echo 'if [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_vsa.py install
|
||||
|
||||
EXPOSE 22
|
||||
@@ -23,3 +23,4 @@ clean:
|
||||
@$(SPHINXBUILD) -M clean "$(SOURCEDIR)" "$(BUILDDIR)" $(SPHINXOPTS) $(O)
|
||||
rm -rf "$(SOURCEDIR)/getting_started/examples"
|
||||
rm -rf "$(SOURCEDIR)/inference/examples"
|
||||
rm -rf "$(SOURCEDIR)/training/examples"
|
||||
|
||||
@@ -11,5 +11,5 @@ commonmark # Required by sphinx-argparse when using :markdownhelp:
|
||||
|
||||
# packages to install to build the documentation
|
||||
cachetools
|
||||
-f https://download.pytorch.org/whl/cpu
|
||||
# -f https://download.pytorch.org/whl/cpu
|
||||
torch
|
||||
@@ -8,12 +8,14 @@ You can easily use the FastVideo Docker image as a custom container on [RunPod](
|
||||
|
||||
Choose a GPU that supports CUDA 12.4
|
||||
|
||||
Pick 1 or 2 L40S GPU(s)
|
||||
|
||||

|
||||
|
||||
When creating your pod template, use this image:
|
||||
|
||||
```
|
||||
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
|
||||
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)):
|
||||
|
||||
@@ -27,6 +27,15 @@ def fix_case(text: str) -> str:
|
||||
"openai": "OpenAI",
|
||||
"multilora": "MultiLoRA",
|
||||
"mlpspeculator": "MLPSpeculator",
|
||||
"finetune": "Finetune",
|
||||
"distillation": "Distillation",
|
||||
"wan": "Wan",
|
||||
"i2v": "I2V",
|
||||
"t2v": "T2V",
|
||||
"1.3b": "1.3B",
|
||||
"14b": "14B",
|
||||
"480p": "480P",
|
||||
"720p": "720P",
|
||||
r"fp\d+": lambda x: x.group(0).upper(), # e.g. fp16, fp32
|
||||
r"int\d+": lambda x: x.group(0).upper(), # e.g. int8, int16
|
||||
}
|
||||
@@ -161,31 +170,35 @@ class Example:
|
||||
return content
|
||||
|
||||
|
||||
def generate_examples(generate_main_index=False):
|
||||
"""
|
||||
Generate example documentation.
|
||||
|
||||
Args:
|
||||
generate_main_index (bool): Whether to generate the main examples index.
|
||||
If False, only category-specific indices will be generated.
|
||||
"""
|
||||
# Create empty indices with dynamic paths
|
||||
@dataclass
|
||||
class NestedStructure:
|
||||
"""Helper class to manage nested documentation structures for training/distillation."""
|
||||
category: str
|
||||
method: str
|
||||
model: str
|
||||
dataset: str
|
||||
example: Example
|
||||
|
||||
@property
|
||||
def filename(self) -> str:
|
||||
return f"{self.model}_{self.dataset}"
|
||||
|
||||
@property
|
||||
def title(self) -> str:
|
||||
return fix_case(self.dataset.replace('_', ' '))
|
||||
|
||||
@property
|
||||
def description(self) -> str:
|
||||
category_name = self.category.title()
|
||||
return f"{category_name} example using the {self.dataset} dataset with the {self.model} model."
|
||||
|
||||
|
||||
def create_category_indices() -> dict[str, Index]:
|
||||
"""Create category indices with their respective configurations."""
|
||||
main_index_dir = ROOT_DIR / "docs/source/examples"
|
||||
if not main_index_dir.exists():
|
||||
main_index_dir.mkdir(parents=True)
|
||||
|
||||
# Create the main examples index only if requested
|
||||
examples_index = None
|
||||
if generate_main_index:
|
||||
examples_index = Index(
|
||||
path=main_index_dir / "examples_index.md",
|
||||
title="💡 Examples",
|
||||
description=
|
||||
"A collection of examples demonstrating usage of FastVideo.\nAll documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.", # noqa: E501
|
||||
caption="Examples",
|
||||
maxdepth=2)
|
||||
|
||||
# Category indices with dynamic paths based on category names
|
||||
category_indices = {
|
||||
"inference":
|
||||
Index(
|
||||
@@ -193,28 +206,54 @@ def generate_examples(generate_main_index=False):
|
||||
"docs/source/inference/examples/examples_inference_index.md",
|
||||
title="🚀 Examples",
|
||||
description=
|
||||
"Inference examples demonstrate how to use FastVideo in an offline setting, where the model is queried for predictions in batches. We recommend starting with <project:basic.md>.", # noqa: E501
|
||||
"Inference examples demonstrate how to use FastVideo inference. We recommend starting with <project:basic.md>.",
|
||||
caption="Examples",
|
||||
maxdepth=1,
|
||||
),
|
||||
"training":
|
||||
Index(
|
||||
path=ROOT_DIR /
|
||||
"docs/source/training/examples/examples_training_index.md",
|
||||
title="🚀 Examples",
|
||||
description=
|
||||
"Training examples demonstrate how to use FastVideo training.",
|
||||
caption="Examples",
|
||||
maxdepth=3,
|
||||
),
|
||||
"distillation":
|
||||
Index(
|
||||
path=ROOT_DIR /
|
||||
"docs/source/distillation/examples/examples_distillation_index.md",
|
||||
title="🚀 Examples",
|
||||
description=
|
||||
"Distillation examples demonstrate how to use FastVideo distillation.",
|
||||
caption="Examples",
|
||||
maxdepth=3,
|
||||
),
|
||||
}
|
||||
|
||||
# Ensure all category doc directories exist
|
||||
for category, index in category_indices.items():
|
||||
category_dir = index.path.parent
|
||||
if not category_dir.exists():
|
||||
category_dir.mkdir(parents=True)
|
||||
for index in category_indices.values():
|
||||
if not index.path.parent.exists():
|
||||
index.path.parent.mkdir(parents=True)
|
||||
|
||||
return category_indices
|
||||
|
||||
|
||||
def find_examples(category_indices: dict[str, Index],
|
||||
generate_main_index: bool) -> list[Example]:
|
||||
"""Find all examples from the examples directory."""
|
||||
examples = []
|
||||
glob_patterns = ["*.py", "*.md", "*.sh"]
|
||||
|
||||
# Find categorised examples
|
||||
for category in category_indices:
|
||||
print(category)
|
||||
category_dir = EXAMPLE_DIR / category
|
||||
globs = [category_dir.glob(pattern) for pattern in glob_patterns]
|
||||
for path in itertools.chain(*globs):
|
||||
examples.append(Example(path, category))
|
||||
# Find examples in subdirectories
|
||||
for path in category_dir.glob("*/*.md"):
|
||||
# Find examples in subdirectories (recursively)
|
||||
for path in category_dir.glob("**/*.md"):
|
||||
examples.append(Example(path.parent, category))
|
||||
|
||||
# Find uncategorised examples only if we're generating a main index
|
||||
@@ -229,36 +268,173 @@ def generate_examples(generate_main_index=False):
|
||||
continue
|
||||
examples.append(Example(path.parent))
|
||||
|
||||
# Create document directories for each category based on category name and generate files
|
||||
for example in sorted(examples, key=lambda e: e.path.stem):
|
||||
print(example)
|
||||
return examples
|
||||
|
||||
|
||||
def create_nested_structures(
|
||||
examples: list[Example]
|
||||
) -> dict[str, dict[str, dict[str, dict[str, NestedStructure]]]]:
|
||||
"""Create nested structures for training and distillation categories."""
|
||||
nested_structures: dict[str, dict[str, dict[str,
|
||||
dict[str,
|
||||
NestedStructure]]]] = {}
|
||||
|
||||
for example in examples:
|
||||
if example.category not in ["training", "distillation"]:
|
||||
continue
|
||||
|
||||
category_dir = EXAMPLE_DIR / example.category
|
||||
relative_path = example.path.relative_to(category_dir)
|
||||
path_parts = relative_path.parts
|
||||
|
||||
# For nested examples like finetune/wan_i2v_14b_480p/crush_smol
|
||||
if len(path_parts) >= 3:
|
||||
method = path_parts[0] # e.g., "finetune"
|
||||
model = path_parts[1] # e.g., "wan_i2v_14b_480p"
|
||||
dataset = path_parts[2] # e.g., "crush_smol"
|
||||
|
||||
# Initialize nested structure
|
||||
if example.category not in nested_structures:
|
||||
nested_structures[example.category] = {}
|
||||
if method not in nested_structures[example.category]:
|
||||
nested_structures[example.category][method] = {}
|
||||
if model not in nested_structures[example.category][method]:
|
||||
nested_structures[example.category][method][model] = {}
|
||||
|
||||
# Store the nested structure
|
||||
nested_structures[
|
||||
example.category][method][model][dataset] = NestedStructure(
|
||||
category=example.category,
|
||||
method=method,
|
||||
model=model,
|
||||
dataset=dataset,
|
||||
example=example)
|
||||
|
||||
return nested_structures
|
||||
|
||||
|
||||
def generate_flat_examples(examples: list[Example],
|
||||
category_indices: dict[str, Index],
|
||||
examples_index: Index | None,
|
||||
generate_main_index: bool) -> None:
|
||||
"""Generate documentation for flat structure examples (inference, etc.)."""
|
||||
for example in examples:
|
||||
if example.category in ["training", "distillation"]:
|
||||
continue # Skip nested structure examples
|
||||
|
||||
# Determine which index to use for this example
|
||||
if example.category is not None and example.category in category_indices:
|
||||
index = category_indices[example.category]
|
||||
elif generate_main_index:
|
||||
assert examples_index is not None
|
||||
index = examples_index # Default to main index if available
|
||||
index = examples_index
|
||||
else:
|
||||
# Skip examples without a category if no main index
|
||||
print(f"Skipping {example.path} (no category and no main index)")
|
||||
continue
|
||||
|
||||
# Place generated example markdown in the same directory as its index
|
||||
# Generate the example documentation
|
||||
doc_path = index.path.parent / f"{example.path.stem}.md"
|
||||
with open(doc_path, "w+") as f:
|
||||
f.write(example.generate())
|
||||
# Add the example to the index
|
||||
index.documents.append(example.path.stem)
|
||||
|
||||
|
||||
def generate_nested_examples(nested_structures: dict[str, dict[str, dict[
|
||||
str, dict[str, NestedStructure]]]], category_indices: dict[str,
|
||||
Index]) -> None:
|
||||
"""Generate documentation for nested structure examples (training, distillation)."""
|
||||
for category_name in ["training", "distillation"]:
|
||||
if category_name not in category_indices or category_name not in nested_structures:
|
||||
continue
|
||||
|
||||
category_index = category_indices[category_name]
|
||||
category_base_dir = category_index.path.parent
|
||||
|
||||
for method, models in nested_structures[category_name].items():
|
||||
# Create method-level index
|
||||
method_index = Index(path=category_base_dir / f"{method}.md",
|
||||
title=fix_case(method),
|
||||
description=f"Examples using {method}.",
|
||||
caption=f"{fix_case(method)} Examples",
|
||||
maxdepth=2)
|
||||
|
||||
for model, datasets in models.items():
|
||||
# Generate dataset examples using the Example class
|
||||
for dataset, nested_struct in datasets.items():
|
||||
doc_path = category_base_dir / f"{nested_struct.filename}.md"
|
||||
with open(doc_path, "w+") as f:
|
||||
f.write(nested_struct.example.generate())
|
||||
|
||||
# Create model-level index
|
||||
model_index = Index(
|
||||
path=category_base_dir / f"{model}.md",
|
||||
title=fix_case(model.replace('_', ' ')),
|
||||
description=f"Examples for the {model} model.",
|
||||
caption=f"{fix_case(model.replace('_', ' '))} Datasets",
|
||||
maxdepth=1)
|
||||
|
||||
# Add dataset indices to model index
|
||||
for dataset, nested_struct in datasets.items():
|
||||
model_index.documents.append(nested_struct.filename)
|
||||
|
||||
# Write model index
|
||||
with open(model_index.path, "w+") as f:
|
||||
f.write(model_index.generate())
|
||||
|
||||
# Add model to method index
|
||||
method_index.documents.append(model)
|
||||
|
||||
# Write method index
|
||||
with open(method_index.path, "w+") as f:
|
||||
f.write(method_index.generate())
|
||||
|
||||
# Add method to main category index
|
||||
category_index.documents.append(method)
|
||||
|
||||
|
||||
def generate_examples(generate_main_index=False):
|
||||
"""
|
||||
Generate example documentation.
|
||||
|
||||
Args:
|
||||
generate_main_index (bool): Whether to generate the main examples index.
|
||||
If False, only category-specific indices will be generated.
|
||||
"""
|
||||
# Create category indices
|
||||
category_indices = create_category_indices()
|
||||
|
||||
# Create the main examples index only if requested
|
||||
examples_index = None
|
||||
if generate_main_index:
|
||||
main_index_dir = ROOT_DIR / "docs/source/examples"
|
||||
examples_index = Index(
|
||||
path=main_index_dir / "examples_index.md",
|
||||
title="💡 Examples",
|
||||
description=
|
||||
"A collection of examples demonstrating usage of FastVideo.\nAll documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.",
|
||||
caption="Examples",
|
||||
maxdepth=2)
|
||||
|
||||
# Find all examples
|
||||
examples = find_examples(category_indices, generate_main_index)
|
||||
|
||||
# Create nested structures for training and distillation
|
||||
nested_structures = create_nested_structures(examples)
|
||||
|
||||
# Generate flat structure examples (inference, etc.)
|
||||
generate_flat_examples(examples, category_indices, examples_index,
|
||||
generate_main_index)
|
||||
|
||||
# Generate nested structure examples (training, distillation)
|
||||
generate_nested_examples(nested_structures, category_indices)
|
||||
|
||||
# Generate the index files for categories
|
||||
for category_index in category_indices.values():
|
||||
if category_index.documents:
|
||||
# Add to main index if it exists
|
||||
if generate_main_index:
|
||||
if generate_main_index and examples_index:
|
||||
main_index_dir = examples_index.path.parent
|
||||
rel_path = category_index.path.relative_to(
|
||||
main_index_dir.parent)
|
||||
assert examples_index is not None
|
||||
examples_index.documents.insert(
|
||||
0,
|
||||
str(rel_path).replace(".md", ""))
|
||||
|
||||
@@ -1,120 +1,18 @@
|
||||
(fastvideo-installation)=
|
||||
(installation-index)=
|
||||
|
||||
# 🔧 Installation
|
||||
|
||||
FastVideo currently only supports Linux and NVIDIA CUDA GPUs.
|
||||
FastVideo supports the following hardware platforms:
|
||||
|
||||
## Requirements
|
||||
:::{toctree}
|
||||
:maxdepth: 1
|
||||
:hidden:
|
||||
|
||||
- **OS: Linux**
|
||||
- **Python: 3.10-3.12**
|
||||
- **CUDA 12.4**
|
||||
- **At least 1 NVIDIA GPU**
|
||||
|
||||
## Set up using Python
|
||||
### Create a new Python environment
|
||||
|
||||
#### Conda
|
||||
You can create a new python environment using [Conda](https://docs.conda.io/projects/conda/en/stable/user-guide/getting-started.html)
|
||||
##### 1. Install Miniconda (if not already installed)
|
||||
|
||||
```bash
|
||||
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
|
||||
bash Miniconda3-latest-Linux-x86_64.sh
|
||||
source ~/.bashrc
|
||||
```
|
||||
|
||||
##### 2. Create and activate a Conda environment for FastVideo
|
||||
|
||||
```bash
|
||||
# (Recommended) Create a new conda environment.
|
||||
conda create -n fastvideo python=3.12 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
:::{note}
|
||||
[PyTorch has deprecated the conda release channel](https://github.com/pytorch/pytorch/issues/138506). If you use `conda`, please only use it to create Python environment rather than installing packages.
|
||||
installation/gpu
|
||||
installation/mps
|
||||
:::
|
||||
|
||||
#### uv
|
||||
|
||||
:::{tip}
|
||||
We highly recommend using `uv` to install FastVideo. In our experience, `uv` speeds up installation by at least 3x.
|
||||
:::
|
||||
|
||||
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
|
||||
|
||||
```console
|
||||
# (Recommended) Create a new uv environment. Use `--seed` to install `pip` and `setuptools` in the environment.
|
||||
uv venv --python 3.12 --seed
|
||||
source .venv/bin/activate
|
||||
```
|
||||
|
||||
### Installation
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
|
||||
# or if you are using uv
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
Also optionally install flash-attn:
|
||||
|
||||
```bash
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
|
||||
#### 1. Clone the FastVideo repository
|
||||
|
||||
```bash
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
```
|
||||
|
||||
#### 2. Install FastVideo
|
||||
|
||||
Basic installation:
|
||||
|
||||
```bash
|
||||
pip install -e .
|
||||
|
||||
# or if you are using uv
|
||||
uv pip install -e .
|
||||
```
|
||||
|
||||
### Optional Dependencies
|
||||
|
||||
#### Flash Attention
|
||||
|
||||
```bash
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
## Set up using Docker
|
||||
We also have prebuilt docker images with FastVideo dependencies pre-installed:
|
||||
[Docker Images](#docker)
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
If you're planning to contribute to FastVideo please see the following page:
|
||||
[Contributor Guide](#developer-overview)
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
### For Basic Inference
|
||||
- NVIDIA GPU with CUDA 12.4 support
|
||||
|
||||
### For Lora Finetuning
|
||||
- 40GB GPU memory each for 2 GPUs with lora
|
||||
- 30GB GPU memory each for 2 GPUs with CPU offload and lora
|
||||
|
||||
### For Full Finetuning/Distillation
|
||||
- Multiple high-memory GPUs recommended (e.g., H100)
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
If you encounter any issues during installation, please open an issue on our [GitHub repository](https://github.com/hao-ai-lab/FastVideo).
|
||||
|
||||
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg) for additional support.
|
||||
- <project:installation/gpu.md>
|
||||
- NVIDIA CUDA
|
||||
- <project:installation/mps.md>
|
||||
- Apple silicon
|
||||
|
||||
@@ -0,0 +1,118 @@
|
||||
# NVIDIA GPU
|
||||
|
||||
Instructions to install FastVideo for NVIDIA CUDA GPUs.
|
||||
|
||||
## Requirements
|
||||
|
||||
- **OS: Linux or Windows WSL**
|
||||
- **Python: 3.10-3.12**
|
||||
- **CUDA 12.4**
|
||||
- **At least 1 NVIDIA GPU**
|
||||
|
||||
## Set up using Python
|
||||
### Create a new Python environment
|
||||
|
||||
#### Conda
|
||||
You can create a new python environment using [Conda](https://docs.conda.io/projects/conda/en/stable/user-guide/getting-started.html)
|
||||
##### 1. Install Miniconda (if not already installed)
|
||||
|
||||
```bash
|
||||
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
|
||||
bash Miniconda3-latest-Linux-x86_64.sh
|
||||
source ~/.bashrc
|
||||
```
|
||||
|
||||
##### 2. Create and activate a Conda environment for FastVideo
|
||||
|
||||
```bash
|
||||
# (Recommended) Create a new conda environment.
|
||||
conda create -n fastvideo python=3.12 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
:::{note}
|
||||
[PyTorch has deprecated the conda release channel](https://github.com/pytorch/pytorch/issues/138506). If you use `conda`, please only use it to create Python environment rather than installing packages.
|
||||
:::
|
||||
|
||||
#### uv
|
||||
|
||||
:::{tip}
|
||||
We highly recommend using `uv` to install FastVideo. In our experience, `uv` speeds up installation by at least 3x.
|
||||
:::
|
||||
|
||||
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
|
||||
|
||||
```console
|
||||
# (Recommended) Create a new uv environment. Use `--seed` to install `pip` and `setuptools` in the environment.
|
||||
uv venv --python 3.12 --seed
|
||||
source .venv/bin/activate
|
||||
```
|
||||
|
||||
### Installation
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
|
||||
# or if you are using uv
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
Also optionally install flash-attn:
|
||||
|
||||
```bash
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
|
||||
#### 1. Clone the FastVideo repository
|
||||
|
||||
```bash
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
```
|
||||
|
||||
#### 2. Install FastVideo
|
||||
|
||||
Basic installation:
|
||||
|
||||
```bash
|
||||
pip install -e .
|
||||
|
||||
# or if you are using uv
|
||||
uv pip install -e .
|
||||
```
|
||||
|
||||
### Optional Dependencies
|
||||
|
||||
#### Flash Attention
|
||||
|
||||
```bash
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
## Set up using Docker
|
||||
We also have prebuilt docker images with FastVideo dependencies pre-installed:
|
||||
[Docker Images](#docker)
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
If you're planning to contribute to FastVideo please see the following page:
|
||||
[Contributor Guide](#developer-overview)
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
### For Basic Inference
|
||||
- NVIDIA GPU with CUDA 12.4 support
|
||||
|
||||
### For Lora Finetuning
|
||||
- 40GB GPU memory each for 2 GPUs with lora
|
||||
- 30GB GPU memory each for 2 GPUs with CPU offload and lora
|
||||
|
||||
### For Full Finetuning/Distillation
|
||||
- Multiple high-memory GPUs recommended (e.g., H100)
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
If you encounter any issues during installation, please open an issue on our [GitHub repository](https://github.com/hao-ai-lab/FastVideo).
|
||||
|
||||
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ) for additional support.
|
||||
@@ -0,0 +1,101 @@
|
||||
# MPS (Apple Silicon)
|
||||
|
||||
Instructions to install FastVideo for Apple Silicon.
|
||||
|
||||
## Requirements
|
||||
|
||||
- **OS: MacOS**
|
||||
- **Python: 3.12.4**
|
||||
|
||||
## Set up using Python
|
||||
|
||||
### Create a new Python environment
|
||||
|
||||
#### Conda
|
||||
|
||||
You can create a new python environment using [Conda](https://docs.conda.io/projects/conda/en/stable/user-guide/getting-started.html)
|
||||
|
||||
##### 1. Install Miniconda (if not already installed)
|
||||
|
||||
```bash
|
||||
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-MacOSX-arm64.sh
|
||||
bash Miniconda3-latest-MacOSX-arm64.sh
|
||||
source ~/.zshrc
|
||||
```
|
||||
|
||||
##### 2. Create and activate a Conda environment for FastVideo
|
||||
|
||||
```bash
|
||||
# (Recommended) Create a new conda environment.
|
||||
conda create -n fastvideo python=3.12.4 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
:::{note}
|
||||
[PyTorch has deprecated the conda release channel](https://github.com/pytorch/pytorch/issues/138506). If you use `conda`, please only use it to create Python environment rather than installing packages.
|
||||
:::
|
||||
|
||||
#### uv
|
||||
|
||||
:::{tip}
|
||||
We highly recommend using `uv` to install FastVideo. In our experience, `uv` speeds up installation by at least 3x.
|
||||
:::
|
||||
|
||||
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
|
||||
|
||||
```console
|
||||
# (Recommended) Create a new uv environment. Use `--seed` to install `pip` and `setuptools` in the environment.
|
||||
uv venv --python 3.12 --seed
|
||||
source .venv/bin/activate
|
||||
```
|
||||
|
||||
### Dependencies
|
||||
|
||||
```
|
||||
brew install ffmpeg
|
||||
```
|
||||
|
||||
### Installation
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
|
||||
# or if you are using uv
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
|
||||
#### 1. Clone the FastVideo repository
|
||||
|
||||
```bash
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
```
|
||||
|
||||
#### 2. Install FastVideo
|
||||
|
||||
Basic installation:
|
||||
|
||||
```bash
|
||||
pip install -e .
|
||||
|
||||
# or if you are using uv
|
||||
uv pip install -e .
|
||||
```
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
If you're planning to contribute to FastVideo please see the following page:
|
||||
[Contributor Guide](#developer-overview)
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
### For Basic Inference
|
||||
|
||||
- Mac M1, M2, M3, or M4 (at least 32 GB RAM is preferable for high quality video generation)
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
If you encounter any issues during installation, please open an issue on our [GitHub repository](https://github.com/hao-ai-lab/FastVideo).
|
||||
|
||||
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ) for additional support.
|
||||
+14
-3
@@ -63,22 +63,33 @@ getting_started/installation
|
||||
:maxdepth: 1
|
||||
|
||||
inference/inference_quick_start
|
||||
inference/examples/examples_inference_index
|
||||
inference/configuration
|
||||
inference/optimizations
|
||||
inference/comfyui
|
||||
inference/support_matrix
|
||||
inference/examples/examples_inference_index
|
||||
inference/cli
|
||||
inference/add_pipeline
|
||||
inference/v0_inference
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Training
|
||||
:maxdepth: 1
|
||||
|
||||
training/examples/examples_training_index
|
||||
training/data_preprocess
|
||||
training/distillation
|
||||
training/finetune
|
||||
<!-- training/finetune -->
|
||||
:::
|
||||
|
||||
<!-- :::{toctree}
|
||||
:caption: Distillation
|
||||
:maxdepth: 1
|
||||
|
||||
distillation/examples/examples_distillation_index
|
||||
distillation/data_preprocess
|
||||
distillation/dmd -->
|
||||
<!-- training/finetune -->
|
||||
:::
|
||||
|
||||
% What is STA Kernel?
|
||||
|
||||
@@ -12,7 +12,7 @@ This guide explains how to implement a custom diffusion pipeline in FastVideo, l
|
||||
4. **Register Your Pipeline** - Make it discoverable by the framework
|
||||
5. **Configure Your Pipeline** - (Coming soon)
|
||||
|
||||
Need help? Join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg).
|
||||
Need help? Join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).
|
||||
|
||||
## Step 1: Pipeline Modules
|
||||
|
||||
|
||||
@@ -27,7 +27,7 @@ fastvideo generate --help
|
||||
### Hardware Configuration
|
||||
|
||||
- `--num-gpus {NUM_GPUS}`: Number of GPUs to use
|
||||
- `--tp-size {TP_SIZE}`: Tensor parallelism size (Typically should match the number of GPUs)
|
||||
- `--tp-size {TP_SIZE}`: Tensor parallelism size (only for the encoder, should not be larger than 1 if text encoder offload is enabled, as layerwise offload + prefetch is faster)
|
||||
- `--sp-size {SP_SIZE}`: Sequence parallelism size (Typically should match the number of GPUs)
|
||||
|
||||
#### Video Configuration
|
||||
@@ -68,7 +68,7 @@ Example configuration file (config.json):
|
||||
"output_path": "outputs/",
|
||||
"num_gpus": 2,
|
||||
"sp_size": 2,
|
||||
"tp_size": 2,
|
||||
"tp_size": 1,
|
||||
"num_frames": 45,
|
||||
"height": 720,
|
||||
"width": 1280,
|
||||
@@ -102,7 +102,7 @@ prompt: "A beautiful woman in a red dress walking down a street"
|
||||
output_path: "outputs/"
|
||||
num_gpus: 2
|
||||
sp_size: 2
|
||||
tp_size: 2
|
||||
tp_size: 1
|
||||
num_frames: 45
|
||||
height: 720
|
||||
width: 1280
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
# FastVideo + ComfyUI
|
||||
|
||||
FastVideo provides a custom node suite for ComfyUI.
|
||||
|
||||
See this [README](https://github.com/hao-ai-lab/FastVideo/tree/main/comfyui) for instructions.
|
||||
@@ -57,8 +57,9 @@ Run the script with:
|
||||
python example.py
|
||||
```
|
||||
|
||||
The generated video will be saved in the current directory under `my_videos/`.
|
||||
The generated video will be saved in the current directory under `my_videos/`
|
||||
|
||||
More inference example scripts can be found in `scripts/inference/`
|
||||
## Available Models
|
||||
|
||||
Please see the [support matrix](#support-matrix) for the list of supported models and their available optimizations.
|
||||
@@ -79,7 +80,6 @@ def main():
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
sampling_param.num_frames = 107
|
||||
sampling_param.image_strength = 0.8 # How much to preserve the original image (0-1)
|
||||
|
||||
# Generate video based on the image
|
||||
prompt = "A photograph coming to life with gentle movement"
|
||||
@@ -121,4 +121,4 @@ If the generated video doesn't match your prompt:
|
||||
- Learn about using [Optimizations](#inference-optimizations)
|
||||
- See [Examples](../examples/examples_inference_index.md) for more usage scenarios
|
||||
- Join our [Community Discord](https://discord.gg/JA7cksDz86).
|
||||
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg).
|
||||
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).
|
||||
|
||||
@@ -1,74 +0,0 @@
|
||||
(v0-inference)=
|
||||
|
||||
# [Deprecated] V0 Inference
|
||||
The following commands and APIs are deprecated but still supported until V1's API can completely replace all the features in this page.
|
||||
|
||||
## 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
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastMochi-diffusers --local_dir=data/FastMochi-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_mochi_sp.sh
|
||||
```
|
||||
@@ -1,7 +1,7 @@
|
||||
(sta-demo)=
|
||||
|
||||
# 🔍 Demo
|
||||
There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
This is is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
<div style="text-align: center;">
|
||||
<video controls width="800">
|
||||
@@ -9,3 +9,9 @@ There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
Your browser does not support the video tag.
|
||||
</video>
|
||||
</div>
|
||||
|
||||
You can run STA using the following command:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_STA.sh
|
||||
```
|
||||
|
||||
@@ -7,70 +7,40 @@ To save GPU memory, we precompute text embeddings and VAE latents to eliminate t
|
||||
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
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/mini_i2v_dataset --local_dir=data/mini_i2v_dataset --repo_type=dataset
|
||||
```
|
||||
|
||||
The folder `crush-smol_raw/` contains raw videos and captions for testing preprocessing, while `crush-smol_preprocessed/` contains latents prepared for testing training.
|
||||
|
||||
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
|
||||
bash scripts/preprocess/v1_preprocess_wan_data_t2v # for wan
|
||||
```
|
||||
|
||||
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:
|
||||
If you wish to create your own dataset for finetuning or distillation, please refer `mini_i2v_dataset/crush-smol_raw/` to structure you video dataset in the following format:
|
||||
|
||||
```
|
||||
path_to_dataset_folder/
|
||||
├── media/
|
||||
│ ├── 0.jpg
|
||||
path_to_your_dataset_folder/
|
||||
├── videos/
|
||||
│ ├── 0.mp4
|
||||
│ ├── 1.mp4
|
||||
│ ├── 2.jpg
|
||||
├── video2caption.json
|
||||
└── merge.txt
|
||||
├── videos.txt
|
||||
└── prompt.txt
|
||||
```
|
||||
|
||||
Format the JSON file as a list, where each item represents a media source:
|
||||
To geranate the `videos2caption.json` and `merge.txt`, run
|
||||
|
||||
For image media,
|
||||
|
||||
```
|
||||
{
|
||||
"path": "0.jpg",
|
||||
"cap": ["captions"]
|
||||
}
|
||||
``` python
|
||||
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
|
||||
```
|
||||
|
||||
For video media,
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/v1_preprocess_****.sh` accordingly and run:
|
||||
|
||||
```
|
||||
{
|
||||
"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
|
||||
bash scripts/preprocess/v1_preprocess_****.sh
|
||||
```
|
||||
|
||||
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
|
||||
|
||||
@@ -16,6 +16,13 @@ bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
## ⚡ Finetune with VSA
|
||||
Follow [data_preprocess.md](#v0-data-preprocess) to get parquet files for preproccessed latent, and then run:
|
||||
|
||||
```bash
|
||||
bash scripts/finetune/finetune_v1_VSA.sh
|
||||
```
|
||||
|
||||
## ⚡ Lora Finetune
|
||||
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
# DMD Distillation Wan2.1-I2V-14B-480P Crush-Smol Example
|
||||
|
||||
Coming soon!
|
||||
@@ -2,7 +2,7 @@ from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -10,8 +10,10 @@ def main():
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
# if num_gpus > 1, FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=True,
|
||||
use_cpu_offload=False
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
@@ -23,7 +25,7 @@ def main():
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
video = generator.generate_video(prompt)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
@@ -34,7 +36,7 @@ def main():
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
video2 = generator.generate_video(prompt2)
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
from fastvideo import VideoGenerator, PipelineConfig
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
def main():
|
||||
config = PipelineConfig.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
config.text_encoder_precisions = ["fp16"]
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
pipeline_config=config,
|
||||
use_fsdp_inference=False, # Disable FSDP for MPS
|
||||
use_cpu_offload=True,
|
||||
text_encoder_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
disable_autocast=False,
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
# Create sampling parameters with reduced number of frames
|
||||
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
sampling_param.num_frames = 3 # Reduce from default 81 to 25 frames bc we have to use the SDPA attn backend for mps
|
||||
sampling_param.height = 256
|
||||
sampling_param.width = 256
|
||||
|
||||
prompt = ("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.")
|
||||
|
||||
video = generator.generate_video(prompt, sampling_param=sampling_param)
|
||||
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, sampling_param=sampling_param)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
|
||||
def main():
|
||||
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
from fastvideo import VideoGenerator, PipelineConfig, SamplingParam
|
||||
|
||||
# from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_fp16"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
model = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
pipeline_config = PipelineConfig.from_pretrained(model)
|
||||
pipeline_config.text_encoder_precisions = ("bf16", )
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model,
|
||||
# if num_gpus > 1, FastVideo will automatically handle distributed setup
|
||||
pipeline_config=pipeline_config,
|
||||
use_fsdp_inference=False, # Disable FSDP for MPS
|
||||
use_cpu_offload=True,
|
||||
text_encoder_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
disable_autocast=False,
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model)
|
||||
sampling_param.num_frames = 30
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = "Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting."
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,55 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "./lora_out"
|
||||
def main():
|
||||
# Initialize VideoGenerator with the Wan model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1,
|
||||
lora_path="benjamin-paine/steamboat-willie-1.3b",
|
||||
lora_nickname="steamboat"
|
||||
)
|
||||
kwargs = {
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 81,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 32,
|
||||
"seed": 42,
|
||||
}
|
||||
# Generate video with LoRA style
|
||||
prompt = "steamboat willie style, golden era animation, close-up of a short fluffy monster kneeling beside a melting red candle. the mood is one of wonder and curiosity, as the monster gazes at the flame with wide eyes and open mouth. Its pose and expression convey a sense of innocence and playfulness, as if it is exploring the world around it for the first time. The use of warm colors and dramatic lighting further enhances the cozy atmosphere of the image."
|
||||
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
# sampling_param=sampling_param,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
negative_prompt=negative_prompt,
|
||||
**kwargs
|
||||
)
|
||||
del generator
|
||||
|
||||
# Until FSDP resharding bug is fixed, multi-lora requires reloading the model or disabling FSDP
|
||||
# see https://github.com/pytorch/pytorch/issues/157209
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1,
|
||||
lora_path="motimalu/wan-flat-color-1.3b-v2",
|
||||
lora_nickname="flat_color"
|
||||
)
|
||||
# generator.set_lora_adapter(lora_nickname="flat_color", lora_path="motimalu/wan-flat-color-1.3b-v2")
|
||||
prompt = "flat color, no lineart, blending, negative space, artist:[john kafka|ponsuke kaikai|hara id 21|yoneyama mai|fuzichoco], 1girl, sakura miko, pink hair, cowboy shot, white shirt, floral print, off shoulder, outdoors, cherry blossom, tree shade, wariza, looking up, falling petals, half-closed eyes, white sky, clouds, live2d animation, upper body, high quality cinematic video of a woman sitting under a sakura tree. Dreamy and lonely, the camera close-ups on the face of the woman as she turns towards the viewer. The Camera is steady, This is a cowboy shot. The animation is smooth and fluid."
|
||||
negative_prompt = "bad quality video,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
negative_prompt=negative_prompt,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,35 @@
|
||||
"""
|
||||
Inference using a LoRA checkpoint from FastVideo trainer.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "./lora_out"
|
||||
def main():
|
||||
# Initialize VideoGenerator with the Wan model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1,
|
||||
lora_path="checkpoints/wan_t2v_finetune_lora/checkpoint-1250/transformer",
|
||||
lora_nickname="crush_smol"
|
||||
)
|
||||
kwargs = {
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 50,
|
||||
"seed": 42,
|
||||
}
|
||||
# Generate video with LoRA style
|
||||
prompt = "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table."
|
||||
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,5 @@
|
||||
# STA Mask Search Examples
|
||||
|
||||
```bash
|
||||
bash examples/inference/sta_mask_search/inference_wan_sta.sh
|
||||
```
|
||||
@@ -0,0 +1,39 @@
|
||||
#!/bin/bash
|
||||
|
||||
export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_wan.json
|
||||
export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
|
||||
export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
|
||||
|
||||
base_port=29503
|
||||
num_gpu=1
|
||||
gpu_ids=$(seq 0 $((num_gpu-1)))
|
||||
skip_time_steps=12
|
||||
|
||||
output_path="inference_results/sta/mask_search_full"
|
||||
STA_mode="STA_searching"
|
||||
for i in $gpu_ids; do
|
||||
port=$((base_port+i))
|
||||
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
|
||||
--prompt_path ./assets/prompt_${i}.txt \
|
||||
--output_path $output_path \
|
||||
--STA_mode $STA_mode &
|
||||
sleep 1
|
||||
done
|
||||
wait
|
||||
echo "STA searching completed"
|
||||
|
||||
output_path="inference_results/sta/mask_search_sparse"
|
||||
STA_mode="STA_tuning"
|
||||
for i in $gpu_ids; do
|
||||
port=$((base_port+i))
|
||||
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
|
||||
--prompt_path ./assets/prompt_${i}.txt \
|
||||
--output_path $output_path \
|
||||
--STA_mode $STA_mode \
|
||||
--skip_time_steps $skip_time_steps &
|
||||
sleep 1
|
||||
done
|
||||
wait
|
||||
echo "STA tuning completed"
|
||||
|
||||
echo "All jobs completed"
|
||||
@@ -0,0 +1,63 @@
|
||||
import os
|
||||
import argparse
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
|
||||
def main(args):
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
# Create a video generator with a pre-trained model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
|
||||
num_gpus=args.num_gpus, # Adjust based on your hardware
|
||||
STA_mode=args.STA_mode,
|
||||
skip_time_steps=args.skip_time_steps
|
||||
)
|
||||
|
||||
# Prompts for your video
|
||||
prompt = args.prompt
|
||||
prompt_path = args.prompt_path
|
||||
negative_prompt = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
|
||||
|
||||
if prompt_path is not None:
|
||||
with open(prompt_path, "r") as f:
|
||||
prompts = f.readlines()
|
||||
else:
|
||||
prompts = [prompt]
|
||||
|
||||
params = SamplingParam(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
fps=args.fps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
return_frames=True, # Also return frames from this call (defaults to False)
|
||||
output_path=args.output_path, # Controls where videos are saved
|
||||
save_video=True,
|
||||
negative_prompt=negative_prompt
|
||||
)
|
||||
|
||||
# Generate the video
|
||||
for prompt in prompts:
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
sampling_param=params,
|
||||
)
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--prompt", type=str, default="A man is dancing.")
|
||||
parser.add_argument("--prompt_path", type=str, default=None)
|
||||
parser.add_argument("--height", type=int, default=768)
|
||||
parser.add_argument("--width", type=int, default=1280)
|
||||
parser.add_argument("--num_frames", type=int, default=69)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=50)
|
||||
parser.add_argument("--fps", type=int, default=16)
|
||||
parser.add_argument("--guidance_scale", type=float, default=5.0)
|
||||
parser.add_argument("--seed", type=int, default=12345)
|
||||
parser.add_argument("--output_path", type=str, default="my_videos/")
|
||||
parser.add_argument("--num_gpus", type=int, default=1)
|
||||
parser.add_argument("--STA_mode", type=str, default="STA_searching")
|
||||
parser.add_argument("--skip_time_steps", type=int, default=12)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,16 @@
|
||||
# Wan2.1-I2V-1.3B-InP Crush-Smol Example
|
||||
These are e2e example scripts for finetuning Wan2.1 T2V 1.3B InP on the crush-smol dataset.
|
||||
|
||||
## Execute the following commands from `FastVideo/` to run training:
|
||||
|
||||
### Download crush-smol dataset:
|
||||
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/download_dataset.sh`
|
||||
|
||||
### Preprocess the videos and captions into latents:
|
||||
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/preprocess_wan_data_i2v.sh`
|
||||
|
||||
### Edit the following file and run finetuning:
|
||||
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/finetune_i2v.sh`
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,94 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
|
||||
NUM_GPUS=8
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_i2v_finetune"
|
||||
--output_dir "$DATA_DIR/outputs/wan_i2v_finetune"
|
||||
--max_train_steps 2000
|
||||
--train_batch_size 4
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 4
|
||||
--tp_size 4
|
||||
--hsdp_replicate_dim 2
|
||||
--hsdp_shard_dim 4
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 2000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,130 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=i2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=4
|
||||
#SBATCH --ntasks=4
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=i2v_output/i2v_%j.out
|
||||
#SBATCH --error=i2v_output/i2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
|
||||
MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_i2v_1_3b_inp/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_i2v_finetune
|
||||
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune_2n"
|
||||
--max_train_steps=2000
|
||||
--train_batch_size=2
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps=1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size $NUM_GPUS
|
||||
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 10
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate=1e-5
|
||||
--mixed_precision="bf16"
|
||||
--checkpointing_steps=1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,24 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_i2v_1_3b_inp/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
--seed 42 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 77 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 16 \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "i2v"
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -0,0 +1,16 @@
|
||||
# Wan2.1-I2V-14B-480P Crush-Smol Example
|
||||
These are e2e examples scripts for finetuning Wan2.1 I2V 14B 480P on the crush-smol dataset.
|
||||
|
||||
## Execute the following commands from `FastVideo/` to run training:
|
||||
|
||||
### Download crush-smol dataset:
|
||||
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/download_dataset.sh`
|
||||
|
||||
### Preprocess the videos and captions into latents:
|
||||
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/preprocess_wan_data_i2v.sh`
|
||||
|
||||
### Edit the following file and run finetuning:
|
||||
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/finetune_i2v.sh`
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,94 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_i2v/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation.json"
|
||||
NUM_GPUS=8
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_i2v_finetune"
|
||||
--output_dir "checkpoints/wan_i2v_finetune"
|
||||
--max_train_steps 2000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 8
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,130 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=i2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=4
|
||||
#SBATCH --ntasks=4
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=i2v_output/i2v_%j.out
|
||||
#SBATCH --error=i2v_output/i2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_i2v/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_i2v_finetune
|
||||
--output_dir="checkpoints/wan_i2v_finetune"
|
||||
--max_train_steps=2000
|
||||
--train_batch_size=2
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps=1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 10
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate=1e-5
|
||||
--mixed_precision="bf16"
|
||||
--checkpointing_steps=1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user