Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
da5ca94091 | ||
|
|
d3ceb67e66 | ||
|
|
7ac153a5ca | ||
|
|
d1e7aa0abd | ||
|
|
2d846c55a1 | ||
|
|
b318063c0a | ||
|
|
4aa307be55 | ||
|
|
055e52e5ea | ||
|
|
7d2069596b | ||
|
|
c45009c9a4 | ||
|
|
b91020b407 | ||
|
|
2dcc5ea4f6 | ||
|
|
359151d9a0 | ||
|
|
ce67cd3729 | ||
|
|
7c554e5da8 | ||
|
|
663ea33ff1 | ||
|
|
3ef04f1654 | ||
|
|
0eced76a41 | ||
|
|
3ab6470d1a | ||
|
|
989a03532c | ||
|
|
fa15369a02 | ||
|
|
78a9cb88d8 | ||
|
|
a0bff12746 | ||
|
|
98f2af94e5 | ||
|
|
46f7b6d574 | ||
|
|
911a6a6a35 | ||
|
|
38c7949d5c | ||
|
|
7e7a0dba9d | ||
|
|
f62e210ae6 | ||
|
|
6ceb4942a0 | ||
|
|
8cae5e4708 | ||
|
|
2a773fa34e | ||
|
|
5357f63327 | ||
|
|
60f61c8101 | ||
|
|
3d75ba8251 | ||
|
|
6c6bcd914d | ||
|
|
f79b08de81 | ||
|
|
f2bc037fff | ||
|
|
86604a684b | ||
|
|
47bd1e0178 | ||
|
|
c41305ad18 | ||
|
|
98ce9034f0 | ||
|
|
0ceff110da | ||
|
|
1d018acb3e | ||
|
|
7d8cf38dbe | ||
|
|
8d483fe4aa | ||
|
|
c1191250bf | ||
|
|
4b7266349a | ||
|
|
22f9b7681f | ||
|
|
589d32cc39 | ||
|
|
89199837db | ||
|
|
d6ebaf1b49 | ||
|
|
fac927777c | ||
|
|
ecbd697dae | ||
|
|
7d4acef64d | ||
|
|
9f0ce517cf | ||
|
|
c718e56b0d | ||
|
|
b65f0316d1 | ||
|
|
8d8bcb76b0 | ||
|
|
5f42748ed1 | ||
|
|
c9005045dc | ||
|
|
6c81befc87 | ||
|
|
dfe0b288e1 | ||
|
|
31200fbb83 | ||
|
|
9185978c55 | ||
|
|
fcba463553 | ||
|
|
2c53d3eecf | ||
|
|
516ecd374a | ||
|
|
3b1b54a74d | ||
|
|
6914e7c904 | ||
|
|
5452369749 | ||
|
|
a113311e77 | ||
|
|
44da97da92 | ||
|
|
f759980a58 | ||
|
|
37e0f8c236 | ||
|
|
51711d5906 | ||
|
|
6375223b16 | ||
|
|
4cb046768d | ||
|
|
3322542444 | ||
|
|
65f707354b | ||
|
|
109e2e7e9d | ||
|
|
cbc3a6bb9d | ||
|
|
2fa8d4ae6d | ||
|
|
7b6c8aee99 | ||
|
|
6284eaa363 | ||
|
|
636524e87f | ||
|
|
202b2f3972 | ||
|
|
247fe273d8 | ||
|
|
cb320dfa3a | ||
|
|
d8bb5abc46 | ||
|
|
cc703eca51 | ||
|
|
81c9df629c | ||
|
|
d3c0c52208 | ||
|
|
744e0555c0 | ||
|
|
3a38f7dfdc | ||
|
|
f572319bd9 | ||
|
|
48528f468c | ||
|
|
4264a80ca9 | ||
|
|
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 |
@@ -1,66 +1,178 @@
|
||||
env:
|
||||
IMAGE_VERSION: "py3.12-latest"
|
||||
BUILDKITE_CLEAN_CHECKOUT: true
|
||||
|
||||
steps:
|
||||
- block: "Start Build"
|
||||
blocked_state: "running"
|
||||
prompt: "Approve build?"
|
||||
- label: "pre-commit"
|
||||
command: ".buildkite/scripts/pre_commit.sh"
|
||||
agents:
|
||||
queue: "default"
|
||||
|
||||
- wait
|
||||
|
||||
- label: "Trigger Tests"
|
||||
command: |
|
||||
echo "Current working directory: $(pwd)"
|
||||
echo "Current branch:"
|
||||
git branch --show-current
|
||||
echo "Full diff:"
|
||||
git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD
|
||||
plugins:
|
||||
- monorepo-diff#v1.4.0:
|
||||
diff: "git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD"
|
||||
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/loaders/**"
|
||||
- "fastvideo/v1/tests/encoders/**"
|
||||
- "fastvideo/models/encoders/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/encoders/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Encoder Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=encoder
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/models/vaes/**"
|
||||
- "fastvideo/v1/models/loaders/**"
|
||||
- "fastvideo/v1/tests/vaes/**"
|
||||
- "fastvideo/models/vaes/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/vaes/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "VAE Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=vae
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/models/dits/**"
|
||||
- "fastvideo/v1/models/loaders/**"
|
||||
- "fastvideo/v1/tests/transformers/**"
|
||||
- "fastvideo/v1/layers/**"
|
||||
- "fastvideo/v1/attention/**"
|
||||
- "fastvideo/models/dits/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/transformers/**"
|
||||
- "fastvideo/layers/**"
|
||||
- "fastvideo/attention/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Transformer Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=transformer
|
||||
agents:
|
||||
queue: "default"
|
||||
- path: "fastvideo/v1/**/*.py"
|
||||
- path:
|
||||
- "fastvideo/**/*.py"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 60m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 45m .buildkite/scripts/pr_test.sh"
|
||||
label: "SSIM Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=ssim
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/tests/lora/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/transformers/**"
|
||||
- "fastvideo/pipelines/**"
|
||||
- "fastvideo/layers/lora/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "LoRA Inference Tests"
|
||||
env:
|
||||
- TEST_TYPE=inference_lora
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/training/*distillation_pipeline.py"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Distillation DMDTests"
|
||||
env:
|
||||
- TEST_TYPE=distillation_dmd
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "LoRA Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training_lora
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "csrc/attn/video_sparse_attn/**"
|
||||
- "csrc/attn/video_sparse_attn/tk/**"
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
- "csrc/attn/video_sparse_attn/config_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/vsa.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Training Tests VSA"
|
||||
env:
|
||||
- TEST_TYPE=training_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/**"
|
||||
- "csrc/attn/sliding_tile_attn/**"
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
- "csrc/attn/sliding_tile_attn/config_sta.py"
|
||||
- "csrc/attn/sliding_tile_attn/st_attn.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Inference Tests STA"
|
||||
env:
|
||||
- TEST_TYPE=inference_sta
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/sliding_tile_attn/**"
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
- "csrc/attn/sliding_tile_attn/config_sta.py"
|
||||
- "csrc/attn/sliding_tile_attn/st_attn.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests STA"
|
||||
env:
|
||||
- TEST_TYPE=precision_sta
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/video_sparse_attn/**"
|
||||
- "csrc/attn/video_sparse_attn/tk/**"
|
||||
- "csrc/attn/tests/test_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
- "csrc/attn/video_sparse_attn/config_vsa.py"
|
||||
- "csrc/attn/video_sparse_attn/vsa.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests VSA"
|
||||
env:
|
||||
- TEST_TYPE=precision_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
|
||||
@@ -31,6 +31,10 @@ 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
|
||||
@@ -46,7 +50,7 @@ else
|
||||
exit 1
|
||||
fi
|
||||
|
||||
MODAL_TEST_FILE="fastvideo/v1/tests/modal/pr_test.py"
|
||||
MODAL_TEST_FILE="fastvideo/tests/modal/pr_test.py"
|
||||
|
||||
if [ -z "${TEST_TYPE:-}" ]; then
|
||||
log "Error: TEST_TYPE environment variable is not set"
|
||||
@@ -54,22 +58,56 @@ if [ -z "${TEST_TYPE:-}" ]; then
|
||||
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="python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
|
||||
;;
|
||||
"vae")
|
||||
log "Running VAE tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
|
||||
;;
|
||||
"transformer")
|
||||
log "Running transformer tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
|
||||
;;
|
||||
"ssim")
|
||||
log "Running SSIM tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
;;
|
||||
"training")
|
||||
log "Running training tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests"
|
||||
;;
|
||||
"training_lora")
|
||||
log "Running LoRA training tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_lora_tests"
|
||||
;;
|
||||
"training_vsa")
|
||||
log "Running training VSA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests_VSA"
|
||||
;;
|
||||
"inference_sta")
|
||||
log "Running inference STA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_STA"
|
||||
;;
|
||||
"precision_sta")
|
||||
log "Running precision STA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_STA"
|
||||
;;
|
||||
"precision_vsa")
|
||||
log "Running precision VSA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_VSA"
|
||||
;;
|
||||
"inference_lora")
|
||||
log "Running LoRA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_lora_tests"
|
||||
;;
|
||||
"distillation_dmd")
|
||||
log "Running distillation DMD tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_distill_dmd_tests"
|
||||
;;
|
||||
*)
|
||||
log "Error: Unknown test type: $TEST_TYPE"
|
||||
|
||||
@@ -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
|
||||
@@ -23,7 +23,7 @@ body:
|
||||
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.
|
||||
Please share your environment with us. You can run the command **python collect_env.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
@@ -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
|
||||
@@ -18,6 +18,12 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
python_3_12_cuda_12_9:
|
||||
description: 'Build Python 3.12 image Cuda 12.9'
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
@@ -49,4 +55,13 @@ jobs:
|
||||
python_version: '3.12'
|
||||
dockerfile_path: docker/Dockerfile.python3.12
|
||||
tag_suffix: py3.12
|
||||
secrets: inherit
|
||||
|
||||
build-python-3-12-cuda-12-9:
|
||||
if: ${{ github.event.inputs.python_3_12_cuda_12_9 == 'true' }}
|
||||
uses: ./.github/workflows/build-image-template.yml
|
||||
with:
|
||||
python_version: '3.12'
|
||||
dockerfile_path: docker/Dockerfile.python3.12.cuda12.9.1
|
||||
tag_suffix: py3.12-cuda12.9.1
|
||||
secrets: inherit
|
||||
@@ -8,14 +8,14 @@ on:
|
||||
- main
|
||||
paths:
|
||||
- "docs/**/*.md"
|
||||
- "fastvideo/v1/examples/**/*.py"
|
||||
- "fastvideo/examples/**/*.py"
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
types: [opened, ready_for_review, synchronize, reopened]
|
||||
paths:
|
||||
- "docs/**/*.md"
|
||||
- "fastvideo/v1/examples/**/*.py"
|
||||
- "fastvideo/examples/**/*.py"
|
||||
|
||||
# Allows you to run this workflow manually from the Actions tab
|
||||
workflow_dispatch:
|
||||
|
||||
@@ -13,4 +13,4 @@
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
@@ -14,13 +14,9 @@ on:
|
||||
- ".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:py3.12-latest)"
|
||||
required: false
|
||||
default: "fastvideo-dev:py3.12-latest"
|
||||
type: string
|
||||
run_encoder_test:
|
||||
description: "Run encoder-test"
|
||||
required: false
|
||||
@@ -56,6 +52,16 @@ on:
|
||||
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
|
||||
@@ -65,6 +71,7 @@ on:
|
||||
env:
|
||||
PYTHONUNBUFFERED: "1"
|
||||
|
||||
|
||||
concurrency:
|
||||
group: pr-test-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
@@ -84,44 +91,70 @@ jobs:
|
||||
training-test: ${{ steps.filter.outputs.training-test }}
|
||||
training-test-VSA: ${{ steps.filter.outputs.training-test-VSA }}
|
||||
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
|
||||
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
|
||||
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
id: filter
|
||||
with:
|
||||
filters: |
|
||||
# Define reusable path patterns
|
||||
common-paths: &common-paths
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
sta-kernel-paths: &sta-kernel-paths
|
||||
- 'csrc/attn/sliding_tile_attn/**'
|
||||
- 'csrc/attn/sliding_tile_attn/tk/**'
|
||||
- 'csrc/attn/sliding_tile_attn/setup.py'
|
||||
- 'csrc/attn/sliding_tile_attn/config_sta.py'
|
||||
- 'csrc/attn/sliding_tile_attn/st_attn.cpp'
|
||||
vsa-kernel-paths: &vsa-kernel-paths
|
||||
- 'csrc/attn/video_sparse_attn/**'
|
||||
- 'csrc/attn/video_sparse_attn/tk/**'
|
||||
- 'csrc/attn/video_sparse_attn/setup.py'
|
||||
- 'csrc/attn/video_sparse_attn/config_vsa.py'
|
||||
- 'csrc/attn/video_sparse_attn/vsa.cpp'
|
||||
vsa-paths: &vsa-paths
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
|
||||
# Actual tests
|
||||
encoder-test:
|
||||
- 'fastvideo/v1/models/encoders/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/encoders/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
- 'fastvideo/models/encoders/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/encoders/**'
|
||||
- *common-paths
|
||||
vae-test:
|
||||
- 'fastvideo/v1/models/vaes/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/vaes/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
- 'fastvideo/models/vaes/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/vaes/**'
|
||||
- *common-paths
|
||||
transformer-test:
|
||||
- 'fastvideo/v1/models/dits/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/transformers/**'
|
||||
- 'fastvideo/v1/layers/**'
|
||||
- 'fastvideo/v1/attention/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
- 'fastvideo/models/dits/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/transformers/**'
|
||||
- 'fastvideo/layers/**'
|
||||
- 'fastvideo/attention/**'
|
||||
- *common-paths
|
||||
training-test:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
training-test-VSA:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
inference-test-STA:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
- *sta-kernel-paths
|
||||
precision-test-STA:
|
||||
- *common-paths
|
||||
- *sta-kernel-paths
|
||||
precision-test-VSA:
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
@@ -134,8 +167,8 @@ jobs:
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
|
||||
test_command: "uv 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/tests/encoders -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -152,8 +185,8 @@ jobs:
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
|
||||
test_command: "uv 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/tests/vaes -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -170,8 +203,8 @@ jobs:
|
||||
gpu_type: "NVIDIA L40S"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
|
||||
test_command: "uv 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/tests/transformers -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -180,8 +213,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:
|
||||
@@ -198,16 +230,16 @@ jobs:
|
||||
volume_size: 200
|
||||
disk_size: 200
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:${{ matrix.python-version.tag }}"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/ssim -vs"
|
||||
timeout_minutes: 60
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
|
||||
training-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(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:
|
||||
@@ -216,8 +248,8 @@ jobs:
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/Vanilla -srP"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/training/Vanilla -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -227,17 +259,17 @@ jobs:
|
||||
training-test-VSA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(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: 1
|
||||
gpu_count: 2
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/VSA -srP"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/training/VSA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -247,17 +279,55 @@ jobs:
|
||||
inference-test-STA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.inference-test-STA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_inference_test_STA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "inference-test-STA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 2
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/inference/STA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
precision-test-STA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-STA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_STA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "precision-test-STA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/inference/STA -srP"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_sta.py"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
precision-test-VSA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-VSA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_VSA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "precision-test-VSA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_vsa.py"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -273,8 +343,8 @@ jobs:
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || '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"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -282,7 +352,8 @@ jobs:
|
||||
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:
|
||||
@@ -299,7 +370,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
|
||||
|
||||
@@ -0,0 +1,28 @@
|
||||
name: Publish to Comfy registry
|
||||
on:
|
||||
workflow_dispatch:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- master
|
||||
paths:
|
||||
- "pyproject.toml"
|
||||
|
||||
permissions:
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
publish-node:
|
||||
name: Publish Custom Node to registry
|
||||
runs-on: ubuntu-latest
|
||||
if: ${{ github.repository_owner == 'hao-ai-lab' }}
|
||||
steps:
|
||||
- name: Check out code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
submodules: true
|
||||
- name: Publish Custom Node
|
||||
uses: Comfy-Org/publish-node-action@v1
|
||||
with:
|
||||
## Add your own personal access token to your Github Repository secrets and reference it here.
|
||||
personal_access_token: ${{ secrets.REGISTRY_ACCESS_TOKEN }}
|
||||
@@ -5,7 +5,7 @@ on:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/attn/sliding_tile_attn/setup.py"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
@@ -23,13 +23,13 @@ jobs:
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/attn
|
||||
cd csrc/attn/sliding_tile_attn
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_sta.py)
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./setup_sta.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
@@ -144,13 +144,13 @@ jobs:
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn # Move into the correct folder
|
||||
cd csrc/attn/sliding_tile_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
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/attn
|
||||
cd csrc/attn/sliding_tile_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)
|
||||
@@ -165,7 +165,7 @@ jobs:
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}
|
||||
path: csrc/attn/dist/*.whl
|
||||
path: csrc/attn/sliding_tile_attn/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
@@ -239,11 +239,11 @@ jobs:
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn # Move into the correct folder
|
||||
cd csrc/attn/sliding_tile_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_sta.py sdist --dist-dir=dist
|
||||
python setup.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/attn/dist/
|
||||
packages-dir: csrc/attn/sliding_tile_attn/dist/
|
||||
|
||||
@@ -0,0 +1,257 @@
|
||||
name: Publish Video Sparse Attention Kernel to PyPI on Version Change
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/attn/video_sparse_attn/setup.py"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
check-version-change:
|
||||
runs-on: ubuntu-latest
|
||||
outputs:
|
||||
version-changed: ${{ steps.check-version.outputs.changed }}
|
||||
new-version: ${{ steps.check-version.outputs.new-version }}
|
||||
steps:
|
||||
- name: Checkout code
|
||||
uses: actions/checkout@v4
|
||||
with:
|
||||
fetch-depth: 2
|
||||
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/attn/video_sparse_attn
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
echo "Version changed from $OLD_VERSION to $NEW_VERSION"
|
||||
echo "changed=true" >> $GITHUB_OUTPUT
|
||||
echo "new-version=$NEW_VERSION" >> $GITHUB_OUTPUT
|
||||
else
|
||||
echo "Version did not change"
|
||||
echo "changed=false" >> $GITHUB_OUTPUT
|
||||
fi
|
||||
|
||||
build_wheels:
|
||||
name: Build Wheel
|
||||
needs: check-version-change
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ${{ matrix.os }}
|
||||
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
# Using ubuntu-20.04 instead of 22.04 for more compatibility (glibc). Ideally we'd use the
|
||||
# manylinux docker image, but I haven't figured out how to install CUDA on manylinux.
|
||||
os: [ubuntu-22.04]
|
||||
python-version: ['3.10', '3.11', '3.12', '3.13']
|
||||
# For version reference https://pytorch.org/get-started/previous-versions/
|
||||
torch-cuda:
|
||||
- torch-version: '2.5.1'
|
||||
cuda-version: '12.4.1'
|
||||
torch-cuda-short: 'cu124'
|
||||
- torch-version: '2.6.0'
|
||||
cuda-version: '12.6.3'
|
||||
torch-cuda-short: 'cu126'
|
||||
- torch-version: '2.7.1'
|
||||
cuda-version: '12.8.0'
|
||||
torch-cuda-short: 'cu128'
|
||||
|
||||
steps:
|
||||
- name: Free up disk space
|
||||
run: |
|
||||
echo "Initial disk space:"
|
||||
df -h
|
||||
|
||||
# Remove large directories
|
||||
sudo rm -rf /usr/share/dotnet
|
||||
sudo rm -rf /usr/local/lib/android
|
||||
sudo rm -rf /opt/ghc
|
||||
sudo rm -rf /usr/local/share/boost
|
||||
sudo rm -rf /usr/share/swift
|
||||
sudo rm -rf /usr/local/lib/node_modules
|
||||
sudo rm -rf /usr/local/share/powershell
|
||||
sudo rm -rf /usr/share/rust
|
||||
sudo rm -rf /usr/local/.ghcup
|
||||
|
||||
# Remove cached files
|
||||
sudo rm -rf /var/lib/apt/lists/*
|
||||
sudo rm -rf /var/cache/apt/archives/*
|
||||
|
||||
echo "Disk space after cleanup:"
|
||||
df -h
|
||||
|
||||
- name: Checkout
|
||||
uses: actions/checkout@v4
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Install CUDA ${{ matrix.torch-cuda.cuda-version }}
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: ${{ matrix.torch-cuda.cuda-version }}
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-${{ matrix.torch-cuda.cuda-version }}
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch ${{ matrix.torch-cuda.torch-version }}+cu${{ matrix.torch-cuda.cuda-version }}
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
|
||||
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
|
||||
pip install typing-extensions==4.12.2
|
||||
# We want to figure out the CUDA version to download pytorch
|
||||
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
|
||||
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
|
||||
pip install --no-cache-dir torch==${{ matrix.torch-cuda.torch-version }} --index-url https://download.pytorch.org/whl/${{matrix.torch-cuda.torch-cuda-short}}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn/video_sparse_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/attn/video_sparse_attn
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-cuda.torch-version }} | cut -d. -f1,2)
|
||||
# Get the correct version format
|
||||
tmpname=cu${CUDA_SHORT_VERSION}torch${TORCH_SHORT_VERSION}
|
||||
wheel_name=$(ls dist/*whl | xargs -n 1 basename | sed "s/-/+$tmpname-/2")
|
||||
# Rename with version information
|
||||
ls dist/*whl |xargs -I {} mv {} dist/${wheel_name}
|
||||
echo "wheel_name=${wheel_name}" >> $GITHUB_ENV
|
||||
|
||||
- name: Upload wheel artifact
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}
|
||||
path: csrc/attn/video_sparse_attn/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
name: Publish package
|
||||
needs: [build_wheels, check-version-change]
|
||||
if: ${{ needs.check-version-change.outputs.version-changed == 'true' || github.event_name == 'workflow_dispatch' }}
|
||||
runs-on: ubuntu-22.04
|
||||
permissions:
|
||||
id-token: write # Needed for OIDC Trusted Publishing
|
||||
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
|
||||
- uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: '3.10'
|
||||
|
||||
- name: Install CUDA 12.4.1
|
||||
uses: Jimver/cuda-toolkit@v0.2.21
|
||||
id: cuda-toolkit
|
||||
with:
|
||||
cuda: 12.4.1
|
||||
linux-local-args: '["--toolkit"]'
|
||||
method: 'network'
|
||||
sub-packages: '["nvcc"]'
|
||||
|
||||
- name: Install dependencies (GCC, Clang, CUDA Paths, Git)
|
||||
run: |
|
||||
sudo apt update
|
||||
sudo apt install -y git patchelf gcc-11 g++-11 clang-11
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Allow Git to Access Safe Directory
|
||||
git config --global --add safe.directory /__w/FastVideo/FastVideo
|
||||
|
||||
# Set CUDA environment variables
|
||||
export CUDA_HOME=/usr/local/cuda-12.4.1
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Verify installation
|
||||
gcc --version
|
||||
g++ --version
|
||||
clang-11 --version
|
||||
nvcc --version
|
||||
|
||||
- name: Install PyTorch 2.5.1+cu12.4.1
|
||||
run: |
|
||||
pip install --upgrade pip
|
||||
# With python 3.13 and torch 2.5.1, unless we update typing-extensions, we get error
|
||||
# AttributeError: attribute '__default__' of 'typing.ParamSpec' objects is not writable
|
||||
pip install typing-extensions==4.12.2
|
||||
# We want to figure out the CUDA version to download pytorch
|
||||
# e.g. we can have system CUDA version being 11.7 but if torch==1.12 then we need to download the wheel from cu116
|
||||
# see https://github.com/pytorch/pytorch/blob/main/RELEASE.md#release-compatibility-matrix
|
||||
export TORCH_CUDA_VERSION=124
|
||||
pip install --no-cache-dir torch==2.5.1 --index-url https://download.pytorch.org/whl/cu${TORCH_CUDA_VERSION}
|
||||
nvcc --version
|
||||
python --version
|
||||
python -c "import torch; print('PyTorch:', torch.__version__)"
|
||||
python -c "import torch; print('CUDA:', torch.version.cuda)"
|
||||
python -c "from torch.utils import cpp_extension; print (cpp_extension.CUDA_HOME)"
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/attn/video_sparse_attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/attn/video_sparse_attn/dist/
|
||||
@@ -20,6 +20,7 @@ samples/
|
||||
data/
|
||||
outputs/
|
||||
outputs_video
|
||||
checkpoints/
|
||||
sbatch.sh
|
||||
*.out
|
||||
env
|
||||
@@ -40,6 +41,8 @@ eggs/
|
||||
docs/_build/
|
||||
docs/source/getting_started/examples/
|
||||
docs/source/inference/examples/
|
||||
docs/source/training/examples/
|
||||
docs/source/distillation/examples/
|
||||
|
||||
# VSCode
|
||||
.vscode/
|
||||
@@ -55,7 +58,9 @@ docs/source/inference/examples/
|
||||
*.pkl
|
||||
|
||||
# Reference videos
|
||||
!fastvideo/v1/tests/ssim/reference_videos/**/*.mp4
|
||||
!fastvideo/tests/ssim/reference_videos/**/*.mp4
|
||||
|
||||
# Static images
|
||||
!docs/source/_static/images/**/*.png
|
||||
!comfyui/assets/**/*.png
|
||||
!comfyui/assets/**/*.gif
|
||||
|
||||
@@ -1,3 +1,7 @@
|
||||
[submodule "csrc/attn/tk"]
|
||||
path = csrc/attn/tk
|
||||
[submodule "csrc/attn/video_sparse_attn/tk"]
|
||||
path = csrc/attn/video_sparse_attn/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
[submodule "csrc/attn/sliding_tile_attn/tk"]
|
||||
path = csrc/attn/sliding_tile_attn/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
@@ -3,7 +3,7 @@ default_stages:
|
||||
- manual # Run in CI
|
||||
exclude: |
|
||||
(?x)(
|
||||
fastvideo/v1/third_party/.*|
|
||||
fastvideo/third_party/.*|
|
||||
csrc/.*|
|
||||
assets/.*|
|
||||
tests/.*|
|
||||
@@ -22,6 +22,7 @@ exclude: |
|
||||
examples/.*|
|
||||
.github/workflows/fastvideo-publish.yml|
|
||||
.github/workflows/sta-publish.yml|
|
||||
.github/workflows/vsa-publish.yml|
|
||||
.github/workflows/build-image-template.yml|
|
||||
docs/source/inference/support_matrix.md
|
||||
)
|
||||
@@ -60,7 +61,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", "--disable-error-code", "union-attr", "--disable-error-code", "override" ]
|
||||
additional_dependencies: [types-cachetools, types-setuptools, types-PyYAML, types-requests]
|
||||
- repo: local
|
||||
hooks:
|
||||
@@ -69,7 +70,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/tests/ssim/" | grep -v "^fastvideo/tests/inference/lora/L40S_reference_videos/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
|
||||
language: system
|
||||
always_run: true
|
||||
pass_filenames: false
|
||||
|
||||
@@ -1,37 +1,41 @@
|
||||
<div align="center">
|
||||
<img src=assets/logo.jpg width="30%"/>
|
||||
<img src=assets/logos/logo.svg width="30%"/>
|
||||
</div>
|
||||
|
||||
**FastVideo is a unified framework for accelerated video generation.**
|
||||
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
|
||||
|
||||
It features a clean, consistent API that works across popular video models, making it easier for developers to author new models and incorporate system- or kernel-level optimizations.
|
||||
With FastVideo's optimizations, you can achieve more than 3x inference improvement compared to other systems.
|
||||
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
|
||||
|
||||
<p align="center">
|
||||
| <a href="https://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://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/S7HLCSTh" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
|
||||
<div align="center">
|
||||
<img src=assets/perf.png width="90%"/>
|
||||
<img src=assets/fastwan.png width="90%"/>
|
||||
</div>
|
||||
|
||||
## NEWS
|
||||
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
- ```2025/02/18```: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
|
||||
## Key Features
|
||||
|
||||
FastVideo has the following features:
|
||||
- End-to-end post-training support:
|
||||
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 to achineve >50x denoising speedup
|
||||
- Data preprocessing pipeline for video data
|
||||
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs
|
||||
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs
|
||||
- State-of-the-art performance optimizations for inference
|
||||
- [Video Sparse Attention](https://arxiv.org/pdf/2505.13389)
|
||||
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
|
||||
- [TeaCache](https://arxiv.org/pdf/2411.19108)
|
||||
- [Sage Attention](https://arxiv.org/abs/2410.02367)
|
||||
- Cutting edge models
|
||||
- Wan2.1 T2V, I2V
|
||||
- HunyuanVideo
|
||||
- FastHunyuan: consistency distilled video diffusion models for 8x inference speedup.
|
||||
- StepVideo T2V
|
||||
- Distillation support
|
||||
- Recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
|
||||
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
|
||||
- Scalable training with FSDP, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
- Memory efficient finetuning with LoRA, precomputed latent, and precomputed text embeddings.
|
||||
- Diverse hardware and OS support
|
||||
- Support H100, A100, 4090
|
||||
- Support Linux, Windows, MacOS
|
||||
|
||||
## Getting Started
|
||||
We recommend using an environment manager such as `Conda` to create a clean environment:
|
||||
@@ -47,17 +51,31 @@ pip install fastvideo
|
||||
|
||||
Please see our [docs](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) for more detailed installation instructions.
|
||||
|
||||
## Sparse Distillation
|
||||
For our sparse distillation techniques, please see our [distillation docs](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) and check out our [blog](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
|
||||
|
||||
See below for recipes and datasets:
|
||||
|
||||
| Model | Sparse Distillation | Dataset |
|
||||
|:-------------------------------------------------------------------------------------------: |:---------------------------------------------------------------------------------------------------------------: |:--------------------------------------------------------------------------------------------------------: |
|
||||
| [FastWan2.1-T2V-1.3B](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P) | [FastVideo Synthetic Wan2.1 480P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k) |
|
||||
| [FastWan2.1-T2V-14B-Preview](https://huggingface.co/FastVideo/FastWan2.1-T2V-14B-Diffusers) | Coming soon! | [FastVideo Synthetic Wan2.1 720P](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x768x1280_250k) |
|
||||
| [FastWan2.2-TI2V-5B](https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-Diffusers) | [Recipe](https://github.com/hao-ai-lab/FastVideo/tree/main/examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free) | [FastVideo Synthetic Wan2.2 720P](https://huggingface.co/datasets/FastVideo/Wan2.2-Syn-121x704x1280_32k) |
|
||||
|
||||
## Inference
|
||||
### Generating Your First Video
|
||||
Here's a minimal example to generate a video using the default settings. Create a file called `example.py` with the following code:
|
||||
Here's a minimal example to generate a video using the default settings. Make sure VSA kernels are [installed](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation.html). Create a file called `example.py` with the following code:
|
||||
|
||||
```python
|
||||
import os
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
def main():
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
|
||||
# Create a video generator with a pre-trained model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1, # Adjust based on your hardware
|
||||
)
|
||||
|
||||
@@ -90,60 +108,63 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
|
||||
- [Contribution Guide](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html)
|
||||
|
||||
## Distillation and Finetuning
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/training/distillation.html)
|
||||
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html)
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html)
|
||||
<!-- - [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html) -->
|
||||
|
||||
## 📑 Development Plan
|
||||
|
||||
<!-- - More distillation methods -->
|
||||
<!-- - [ ] Add Distribution Matching Distillation -->
|
||||
- More models support
|
||||
<!-- - [ ] Add CogvideoX model -->
|
||||
- [x] Add StepVideo to V1
|
||||
- Optimization features
|
||||
- [x] Teacache in V1
|
||||
- [x] SageAttention in V1
|
||||
- Code updates
|
||||
- [x] V1 Configuration API
|
||||
- [ ] Support Training in V1
|
||||
More FastWan Models Coming Soon!
|
||||
- [ ] Add FastWan2.1-T2V-14B
|
||||
- [ ] Add FastWan2.2-T2V-14B
|
||||
- [ ] Add FastWan2.2-I2V-14B
|
||||
<!-- - Optimization features
|
||||
- Code updates -->
|
||||
<!-- - [ ] fp8 support -->
|
||||
<!-- - [ ] faster load model and save model support -->
|
||||
|
||||
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
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:
|
||||
- [PCM](https://github.com/G-U-N/Phased-Consistency-Model)
|
||||
- [Wan-Video](https://github.com/Wan-Video)
|
||||
- [ThunderKittens](https://github.com/HazyResearch/ThunderKittens)
|
||||
- [Triton](https://github.com/triton-lang/triton)
|
||||
- [DMD2](https://github.com/tianweiy/DMD2)
|
||||
- [diffusers](https://github.com/huggingface/diffusers)
|
||||
- [OpenSoraPlan](https://github.com/PKU-YuanGroup/Open-Sora-Plan)
|
||||
- [xDiT](https://github.com/xdit-project/xDiT)
|
||||
- [vLLM](https://github.com/vllm-project/vllm)
|
||||
- [SGLang](https://github.com/sgl-project/sglang)
|
||||
|
||||
We thank MBZUAI and [Anyscale](https://www.anyscale.com/) for their support throughout this project.
|
||||
We thank [MBZUAI](https://ifm.mbzuai.ac.ae/), [Anyscale](https://www.anyscale.com/), and [GMI Cloud](https://www.gmicloud.ai/) for their support throughout this project.
|
||||
|
||||
## Citation
|
||||
If you use FastVideo for your research, please cite our paper:
|
||||
If you find FastVideo useful, please considering citing our work:
|
||||
|
||||
```bibtex
|
||||
@misc{zhang2025fastvideogenerationsliding,
|
||||
title={Fast Video Generation with Sliding Tile Attention},
|
||||
author={Peiyuan Zhang and Yongqi Chen and Runlong Su and Hangliang Ding and Ion Stoica and Zhenghong Liu and Hao Zhang},
|
||||
year={2025},
|
||||
eprint={2502.04507},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV},
|
||||
url={https://arxiv.org/abs/2502.04507},
|
||||
@software{fastvideo2024,
|
||||
title = {FastVideo: A Unified Framework for Accelerated Video Generation},
|
||||
author = {The FastVideo Team},
|
||||
url = {https://github.com/hao-ai-lab/FastVideo},
|
||||
month = apr,
|
||||
year = {2024},
|
||||
}
|
||||
@misc{ding2025efficientvditefficientvideodiffusion,
|
||||
title={Efficient-vDiT: Efficient Video Diffusion Transformers With Attention Tile},
|
||||
author={Hangliang Ding and Dacheng Li and Runlong Su and Peiyuan Zhang and Zhijie Deng and Ion Stoica and Hao Zhang},
|
||||
year={2025},
|
||||
eprint={2502.06155},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV},
|
||||
url={https://arxiv.org/abs/2502.06155},
|
||||
|
||||
@article{zhang2025vsa,
|
||||
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
|
||||
author={Zhang, Peiyuan and Huang, Haofeng and Chen, Yongqi and Lin, Will and Liu, Zhengzhong and Stoica, Ion and Xing, Eric and Zhang, Hao},
|
||||
journal={arXiv preprint arXiv:2505.13389},
|
||||
year={2025}
|
||||
}
|
||||
|
||||
@article{zhang2025fast,
|
||||
title={Fast video generation with sliding tile attention},
|
||||
author={Zhang, Peiyuan and Chen, Yongqi and Su, Runlong and Ding, Hangliang and Stoica, Ion and Liu, Zhengzhong and Zhang, Hao},
|
||||
journal={arXiv preprint arXiv:2502.04507},
|
||||
year={2025}
|
||||
}
|
||||
```
|
||||
|
||||
@@ -0,0 +1,15 @@
|
||||
try:
|
||||
from .comfyui.video_generator.nodes import (NODE_CLASS_MAPPINGS,
|
||||
NODE_DISPLAY_NAME_MAPPINGS)
|
||||
WEB_DIRECTORY = "./web"
|
||||
__all__ = [
|
||||
'NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY'
|
||||
]
|
||||
except ImportError:
|
||||
# ComfyUI environment not available, skip comfyui imports
|
||||
NODE_CLASS_MAPPINGS = {}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {}
|
||||
WEB_DIRECTORY = "./web"
|
||||
__all__ = [
|
||||
'NODE_CLASS_MAPPINGS', 'NODE_DISPLAY_NAME_MAPPINGS', 'WEB_DIRECTORY'
|
||||
]
|
||||
|
After Width: | Height: | Size: 194 KiB |
|
Before Width: | Height: | Size: 149 KiB |
@@ -0,0 +1,6 @@
|
||||
<svg width="160" height="93" viewBox="0 0 160 93" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<path d="M28.8511 91.66L57.6319 1.86368H64.5394L35.7585 91.66H28.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
<path d="M15.0376 91.66L43.8185 1.86368H46.1209L17.3401 91.66H15.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
<path d="M1.22217 91.66L30.003 1.86366H31.1543L2.3734 91.66H1.22217Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
|
||||
<path d="M71.4465 1.86483L42.666 91.6599H69.144L78.3538 58.2746H123.251L129.007 39.855H84.1099L89.866 22.5868H152.032L157.788 1.86483H71.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 691 B |
@@ -0,0 +1,18 @@
|
||||
<svg width="252" height="105" viewBox="0 0 252 105" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<path d="M89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457Z" fill="#356CFF"/>
|
||||
<path fill-rule="evenodd" clip-rule="evenodd" d="M96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM100.768 13.1217L87.7028 29.4852H103.143L100.768 13.1217Z" fill="#356CFF"/>
|
||||
<path d="M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057Z" fill="#356CFF"/>
|
||||
<path fill-rule="evenodd" clip-rule="evenodd" d="M108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM109.081 90.697L116.802 65.8487C116.802 65.8487 120.959 65.8487 132.242 65.8487C143.525 65.8487 137.586 78.5759 135.211 84.0304C133.307 88.4021 127.491 90.697 122.74 90.697C117.989 90.697 109.081 90.697 109.081 90.697Z" fill="#356CFF"/>
|
||||
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944H159.747C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273L125.188 48.273L124 37.97L147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852H131.836C120.142 29.4852 125.897 1.00043 141.337 1.00043L173.188 1.00056Z" fill="#356CFF"/>
|
||||
<path d="M179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056Z" fill="#356CFF"/>
|
||||
<path d="M161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457Z" fill="#356CFF"/>
|
||||
<path fill-rule="evenodd" clip-rule="evenodd" d="M230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM237.948 77.9692C239.984 70.6965 240.917 65.242 228.446 65.242C215.975 65.242 211.818 71.9087 210.037 77.9692C208.255 84.0298 208.255 91.3025 219.538 91.3025C230.821 91.3025 235.911 85.2419 237.948 77.9692Z" fill="#356CFF"/>
|
||||
<path d="M173.188 1.00056L168.625 11.9096C168.625 11.9096 149.386 11.9092 142.525 11.9095C135.664 11.9098 136.586 20.3944 141.337 20.3944M173.188 1.00056C173.188 1.00056 156.777 1.00043 141.337 1.00043M173.188 1.00056L141.337 1.00043M141.337 20.3944C146.088 20.3944 150.839 20.3944 159.747 20.3944M141.337 20.3944H159.747M159.747 20.3944C168.654 20.3944 166.961 33.6899 163.904 38.5761C160.467 44.0675 157.371 48.273 148.463 48.273M148.463 48.273C139.556 48.273 125.188 48.273 125.188 48.273M148.463 48.273L125.188 48.273M125.188 48.273L124 37.97M124 37.97C124 37.97 141.931 37.97 147.87 37.97M124 37.97L147.87 37.97M147.87 37.97C153.808 37.97 156.184 29.4852 151.433 29.4852M151.433 29.4852C146.682 29.4852 138.962 29.4852 131.836 29.4852M151.433 29.4852H131.836M131.836 29.4852C120.142 29.4852 125.897 1.00043 141.337 1.00043M37.2252 1.00057L22.3789 48.273H36.0375L40.7884 30.6974L62.6727 30.6974L69.6727 21.0005L43.7576 21.0004L46.7269 11.9096L77.6727 11.9096L86 1.00057L37.2252 1.00057ZM96.0167 1.00057H112.645L118.583 48.273H104.924L103.737 39.7882H79.9827L67.5117 55.5457H85.3273L43.1638 101H28.3174L22.3789 55.5457H33.6621L38.4129 91.3031L58.604 68.2729H44.3515L96.0167 1.00057ZM87.7028 29.4852L100.768 13.1217L103.143 29.4852H87.7028ZM89.4843 55.5457H101.361L87.7028 101H74.638L89.4843 55.5457ZM108.488 55.5457L94.2351 101C94.2351 101 105.518 101 120.959 101C136.399 101 144.078 93.0133 148.276 79.788C152.432 68.0157 153.027 55.5457 136.399 55.5457C119.771 55.5457 108.488 55.5457 108.488 55.5457ZM116.802 65.8487L109.081 90.697C109.081 90.697 117.989 90.697 122.74 90.697C127.491 90.697 133.307 88.4021 135.211 84.0304C137.586 78.5759 143.525 65.8487 132.242 65.8487C120.959 65.8487 116.802 65.8487 116.802 65.8487ZM179.938 1.00056L175.688 11.9096L191.221 11.9096L179.938 48.273H192.409L203.692 11.9096L219.132 11.9095L223.289 1.00043L179.938 1.00056ZM161.341 55.5457H202.845L198.5 65.8487H169.654L167.279 73.7268H188.5L184.749 82.8177H164.31L161.934 90.697H190.251L186.624 101H146.494L161.341 55.5457ZM230.821 54.9391C255.169 54.9391 251.776 67.0602 249.231 77.9692C246.686 88.8783 240.917 101 217.757 101C194.596 101 195.606 88.8783 199.347 77.9692C203.089 67.0602 206.473 54.9391 230.821 54.9391ZM228.446 65.242C240.917 65.242 239.984 70.6965 237.948 77.9692C235.911 85.2419 230.821 91.3025 219.538 91.3025C208.255 91.3025 208.255 84.0298 210.037 77.9692C211.818 71.9087 215.975 65.242 228.446 65.242Z" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M15.2524 55.5451L21.191 100.999L24.7541 100.999L18.8156 55.5451L15.2524 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M8.12646 55.5451L14.065 100.999L15.2527 100.999L9.31417 55.5451L8.12646 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M1 55.5451L6.93853 100.999L7.53239 100.999L1.59385 55.5451L1 55.5451Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
|
||||
<path d="M15.2524 48.2724L30.0988 1H33.6619L18.8156 48.2724H15.2524Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M8.12646 48.2724L22.9728 1H24.1605L9.31417 48.2724H8.12646Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.18771"/>
|
||||
<path d="M1 48.2724L15.8463 1H16.4402L1.59385 48.2724H1Z" fill="#356CFF" stroke="#356CFF" stroke-width="0.593853"/>
|
||||
<path d="M85.3271 55.5457H67.5116L87 12.7363L44.3513 68.2729H58.6038L43.1636 101L85.3271 55.5457Z" fill="#FDC717" stroke="#FDC717" stroke-width="1.18771" stroke-miterlimit="16"/>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 5.7 KiB |
@@ -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"
|
||||
@@ -16,7 +16,7 @@ import sys
|
||||
# Run it with `python collect_env.py` or `python -m torch.utils.collect_env`
|
||||
from collections import namedtuple
|
||||
|
||||
from fastvideo.v1.envs import environment_variables
|
||||
from fastvideo.envs import environment_variables
|
||||
|
||||
try:
|
||||
import torch
|
||||
@@ -62,6 +62,7 @@ SystemEnv = namedtuple(
|
||||
DEFAULT_CONDA_PATTERNS = {
|
||||
"torch",
|
||||
"numpy",
|
||||
"mypy"
|
||||
"cudatoolkit",
|
||||
"soumith",
|
||||
"mkl",
|
||||
@@ -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/configs/wan_14B_i2v_480p_pipeline.json)
|
||||
- [FastHunyuan-diffusers](https://github.com/hao-ai-lab/FastVideo/blob/main/fastvideo/configs/fasthunyuan_t2v.json)
|
||||
|
||||
### Node Configuration
|
||||
|
||||
#### Video Generator
|
||||
|
||||
- **prompt**: Text description of the video to generate
|
||||
- **output_path**: Directory where generated videos will be saved
|
||||
- **num_gpus**: Number of GPUs to use for generation
|
||||
- **model_path**: Path to the FastVideo model
|
||||
- **embedded_cfg_scale**: Classifier-free guidance scale
|
||||
- **sp_size**: Sequence parallelism size (usually should match num_gpus)
|
||||
- **tp_size**: Tensor parallelism size (usually should match num_gpus)
|
||||
- **precision**: Model precision (fp16 or bf16)
|
||||
|
||||
`model_path takes either a model id from huggingface or a local path to a model. Models by default will be downloaded to ~/.cache/huggingface/hub/ and cached for subsequent runs.`
|
||||
|
||||
#### Inference Args
|
||||
|
||||
- **height/width**: Resolution of the output video
|
||||
- **num_frames**: Number of frames to generate
|
||||
- **num_inference_steps**: Number of diffusion steps per frame
|
||||
- **guidance_scale**: Classifier-free guidance scale
|
||||
- **flow_shift**: Frame flow shift parameter
|
||||
- **seed**: Random seed for reproducible generation
|
||||
- **fps**: Frames per second of the output video
|
||||
- **image_path**: Optional path to input image for conditioning (for i2v models)
|
||||
|
||||
## Memory Management
|
||||
|
||||
Models will remain loaded in GPU memory between runs when you only change inference arguments (such as prompt, resolution, frame count, FPS, guidance scale, etc.) or the prompt text. This allows for faster subsequent generations since the model doesn't need to be reloaded.
|
||||
|
||||
However, if you need to change the following parameters, you will need to restart the ComfyUI server:
|
||||
- **Number of GPUs** (`num_gpus`)
|
||||
- **Model path** (`model_path`)
|
||||
- **Tensor parallelism size** (`tp_size`)
|
||||
- **Sequence parallelism size** (`sp_size`)
|
||||
|
||||
These parameters affect the model's distribution across GPUs and require a complete reinitialization of the model pipeline.
|
||||
|
||||
## Example workflows
|
||||
|
||||
### Text to Video
|
||||
|
||||
FastVideo-FastHunyuan-diffusers
|
||||
|
||||

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

|
||||
|
||||
- [Wan2.1-I2V-14B-480P-Diffusers.json](./examples/Wan2.1-I2V-14B-480P-Diffusers.json)
|
||||
|
||||
## License
|
||||
|
||||
This project is licensed under Apache 2.0.
|
||||
@@ -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']
|
||||
|
After Width: | Height: | Size: 1.3 MiB |
@@ -0,0 +1,6 @@
|
||||
<svg width="160" height="93" viewBox="0 0 160 93" fill="none" xmlns="http://www.w3.org/2000/svg">
|
||||
<path d="M28.8511 91.66L57.6319 1.86368H64.5394L35.7585 91.66H28.8511Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
<path d="M15.0376 91.66L43.8185 1.86368H46.1209L17.3401 91.66H15.0376Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
<path d="M1.22217 91.66L30.003 1.86366H31.1543L2.3734 91.66H1.22217Z" fill="#356CFF" stroke="#356CFF" stroke-width="1.15122"/>
|
||||
<path d="M71.4465 1.86483L42.666 91.6599H69.144L78.3538 58.2746H123.251L129.007 39.855H84.1099L89.866 22.5868H152.032L157.788 1.86483H71.4465Z" fill="#356CFF" stroke="#356CFF" stroke-width="2.30244"/>
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 691 B |
|
After Width: | Height: | Size: 8.7 MiB |
|
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"
|
||||
},
|
||||
"dit_cpu_offload": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 3,
|
||||
"type": "VHS_LoadVideoPath",
|
||||
"pos": [
|
||||
1350.136962890625,
|
||||
331.20361328125
|
||||
],
|
||||
"size": [
|
||||
231.8896484375,
|
||||
286
|
||||
],
|
||||
"flags": {},
|
||||
"order": 5,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"shape": 7,
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"shape": 7,
|
||||
"type": "VAE",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "video",
|
||||
"type": "STRING",
|
||||
"widget": {
|
||||
"name": "video"
|
||||
},
|
||||
"link": 4
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "IMAGE",
|
||||
"type": "IMAGE",
|
||||
"links": [
|
||||
5
|
||||
]
|
||||
},
|
||||
{
|
||||
"name": "frame_count",
|
||||
"type": "INT",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"type": "AUDIO",
|
||||
"links": null
|
||||
},
|
||||
{
|
||||
"name": "video_info",
|
||||
"type": "VHS_VIDEOINFO",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VHS_LoadVideoPath"
|
||||
},
|
||||
"widgets_values": {
|
||||
"video": "",
|
||||
"force_rate": 0,
|
||||
"custom_width": 0,
|
||||
"custom_height": 0,
|
||||
"frame_load_cap": 0,
|
||||
"skip_first_frames": 0,
|
||||
"select_every_nth": 1,
|
||||
"format": "Wan",
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"filename": "",
|
||||
"type": "path",
|
||||
"format": "video/",
|
||||
"force_rate": 0,
|
||||
"custom_width": 0,
|
||||
"custom_height": 0,
|
||||
"frame_load_cap": 0,
|
||||
"skip_first_frames": 0,
|
||||
"select_every_nth": 1
|
||||
}
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 2,
|
||||
"type": "InferenceArgs",
|
||||
"pos": [
|
||||
411.46307373046875,
|
||||
178.18182373046875
|
||||
],
|
||||
"size": [
|
||||
278.73828125,
|
||||
298
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "inference_args",
|
||||
"type": "INFERENCE_ARGS",
|
||||
"links": [
|
||||
3
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "InferenceArgs"
|
||||
},
|
||||
"widgets_values": [
|
||||
720,
|
||||
1280,
|
||||
45,
|
||||
6,
|
||||
-99999,
|
||||
-99999,
|
||||
1025,
|
||||
"fixed",
|
||||
24,
|
||||
-99999,
|
||||
-99999
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"height": {
|
||||
"isAuto": false,
|
||||
"value": 720,
|
||||
"cachedValue": 720
|
||||
},
|
||||
"width": {
|
||||
"isAuto": false,
|
||||
"value": 1280,
|
||||
"cachedValue": 1280
|
||||
},
|
||||
"num_frames": {
|
||||
"isAuto": false,
|
||||
"value": 45,
|
||||
"cachedValue": 45
|
||||
},
|
||||
"num_inference_steps": {
|
||||
"isAuto": false,
|
||||
"value": 6,
|
||||
"cachedValue": 6
|
||||
},
|
||||
"guidance_scale": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 1
|
||||
},
|
||||
"flow_shift": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 17
|
||||
},
|
||||
"seed": {
|
||||
"isAuto": false,
|
||||
"value": 1025,
|
||||
"cachedValue": 1024
|
||||
},
|
||||
"fps": {
|
||||
"isAuto": false,
|
||||
"value": 24,
|
||||
"cachedValue": 24
|
||||
},
|
||||
"image_path": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": "X://insert/path/here.mp4"
|
||||
},
|
||||
"enable_teacache": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 8,
|
||||
"type": "VHS_VideoCombine",
|
||||
"pos": [
|
||||
1668.3499755859375,
|
||||
328.22625732421875
|
||||
],
|
||||
"size": [
|
||||
507.507080078125,
|
||||
622.2227172851562
|
||||
],
|
||||
"flags": {},
|
||||
"order": 6,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "images",
|
||||
"type": "IMAGE",
|
||||
"link": 5
|
||||
},
|
||||
{
|
||||
"name": "audio",
|
||||
"shape": 7,
|
||||
"type": "AUDIO",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "meta_batch",
|
||||
"shape": 7,
|
||||
"type": "VHS_BatchManager",
|
||||
"link": null
|
||||
},
|
||||
{
|
||||
"name": "vae",
|
||||
"shape": 7,
|
||||
"type": "VAE",
|
||||
"link": null
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "Filenames",
|
||||
"type": "VHS_FILENAMES",
|
||||
"links": null
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VHS_VideoCombine"
|
||||
},
|
||||
"widgets_values": {
|
||||
"frame_rate": 24,
|
||||
"loop_count": 0,
|
||||
"filename_prefix": "",
|
||||
"format": "video/h264-mp4",
|
||||
"pix_fmt": "yuv420p",
|
||||
"crf": 19,
|
||||
"save_metadata": true,
|
||||
"trim_to_audio": false,
|
||||
"pingpong": false,
|
||||
"save_output": false,
|
||||
"videopreview": {
|
||||
"hidden": false,
|
||||
"paused": false,
|
||||
"params": {
|
||||
"filename": "._00003.mp4",
|
||||
"subfolder": "",
|
||||
"type": "temp",
|
||||
"format": "video/h264-mp4",
|
||||
"frame_rate": 24,
|
||||
"workflow": "._00003.png",
|
||||
"fullpath": "/workspace/ComfyUI/temp/._00003.mp4"
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
2,
|
||||
4,
|
||||
0,
|
||||
1,
|
||||
1,
|
||||
"VAE_CONFIG"
|
||||
],
|
||||
[
|
||||
3,
|
||||
2,
|
||||
0,
|
||||
1,
|
||||
0,
|
||||
"INFERENCE_ARGS"
|
||||
],
|
||||
[
|
||||
4,
|
||||
1,
|
||||
0,
|
||||
3,
|
||||
2,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
5,
|
||||
3,
|
||||
0,
|
||||
8,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
6,
|
||||
6,
|
||||
0,
|
||||
1,
|
||||
3,
|
||||
"DIT_CONFIG"
|
||||
],
|
||||
[
|
||||
7,
|
||||
5,
|
||||
0,
|
||||
1,
|
||||
2,
|
||||
"TEXT_ENCODER_CONFIG"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.9090909090909091,
|
||||
"offset": [
|
||||
112.86678372727341,
|
||||
-71.45635903989245
|
||||
]
|
||||
},
|
||||
"frontendVersion": "1.20.4",
|
||||
"VHS_latentpreview": false,
|
||||
"VHS_latentpreviewrate": 0,
|
||||
"VHS_MetadataImage": true,
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
@@ -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"
|
||||
},
|
||||
"dit_cpu_offload": {
|
||||
"isAuto": true,
|
||||
"value": true,
|
||||
"cachedValue": true
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "TextEncoderConfig",
|
||||
"pos": [
|
||||
416.4937744140625,
|
||||
953.6171875
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
106
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "text_encoder_config",
|
||||
"type": "TEXT_ENCODER_CONFIG",
|
||||
"links": [
|
||||
7
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "TextEncoderConfig"
|
||||
},
|
||||
"widgets_values": [
|
||||
"",
|
||||
"",
|
||||
""
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"prefix": {
|
||||
"isAuto": true,
|
||||
"value": "",
|
||||
"cachedValue": ""
|
||||
},
|
||||
"quant_config": {
|
||||
"isAuto": true,
|
||||
"value": "",
|
||||
"cachedValue": ""
|
||||
},
|
||||
"lora_config": {
|
||||
"isAuto": true,
|
||||
"value": "",
|
||||
"cachedValue": ""
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 6,
|
||||
"type": "DITConfig",
|
||||
"pos": [
|
||||
415.1928405761719,
|
||||
1154.1573486328125
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
82
|
||||
],
|
||||
"flags": {},
|
||||
"order": 3,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "dit_config",
|
||||
"type": "DIT_CONFIG",
|
||||
"links": [
|
||||
6
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "DITConfig"
|
||||
},
|
||||
"widgets_values": [
|
||||
"",
|
||||
""
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"prefix": {
|
||||
"isAuto": true,
|
||||
"value": "",
|
||||
"cachedValue": ""
|
||||
},
|
||||
"quant_config": {
|
||||
"isAuto": true,
|
||||
"value": "",
|
||||
"cachedValue": ""
|
||||
}
|
||||
}
|
||||
}
|
||||
],
|
||||
"links": [
|
||||
[
|
||||
1,
|
||||
7,
|
||||
0,
|
||||
2,
|
||||
0,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
2,
|
||||
4,
|
||||
0,
|
||||
1,
|
||||
1,
|
||||
"VAE_CONFIG"
|
||||
],
|
||||
[
|
||||
3,
|
||||
2,
|
||||
0,
|
||||
1,
|
||||
0,
|
||||
"INFERENCE_ARGS"
|
||||
],
|
||||
[
|
||||
4,
|
||||
1,
|
||||
0,
|
||||
3,
|
||||
2,
|
||||
"STRING"
|
||||
],
|
||||
[
|
||||
5,
|
||||
3,
|
||||
0,
|
||||
8,
|
||||
0,
|
||||
"IMAGE"
|
||||
],
|
||||
[
|
||||
6,
|
||||
6,
|
||||
0,
|
||||
1,
|
||||
3,
|
||||
"DIT_CONFIG"
|
||||
],
|
||||
[
|
||||
7,
|
||||
5,
|
||||
0,
|
||||
1,
|
||||
2,
|
||||
"TEXT_ENCODER_CONFIG"
|
||||
]
|
||||
],
|
||||
"groups": [],
|
||||
"config": {},
|
||||
"extra": {
|
||||
"ds": {
|
||||
"scale": 0.8264462809917354,
|
||||
"offset": [
|
||||
646.7950212991898,
|
||||
66.17259910028655
|
||||
]
|
||||
},
|
||||
"frontendVersion": "1.20.4",
|
||||
"VHS_latentpreview": false,
|
||||
"VHS_latentpreviewrate": 0,
|
||||
"VHS_MetadataImage": true,
|
||||
"VHS_KeepIntermediate": true
|
||||
},
|
||||
"version": 0.4
|
||||
}
|
||||
@@ -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"
|
||||
}),
|
||||
"dit_cpu_offload": ([True, False], {
|
||||
"default": False
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def VALIDATE_INPUTS(cls, **kwargs):
|
||||
return True
|
||||
|
||||
RETURN_TYPES = ("STRING", )
|
||||
RETURN_NAMES = ("video_path", )
|
||||
FUNCTION = "launch_inference"
|
||||
CATEGORY = "fastvideo"
|
||||
|
||||
generator: FastVideoGenerator | None = None
|
||||
_interrupt_thread: threading.Thread | None = None
|
||||
_generation_active: bool = False
|
||||
_generation_interrupted: bool = False
|
||||
_interrupt_event: threading.Event = threading.Event()
|
||||
_generation_thread: threading.Thread | None = None
|
||||
_generation_result: str | None = None
|
||||
_generation_exception: Exception | None = None
|
||||
|
||||
def _monitor_for_interruption(self):
|
||||
"""Background thread that monitors for interruption requests"""
|
||||
time.sleep(2) # Give the generation thread time to send execute_forward
|
||||
|
||||
while self._generation_active and not self._interrupt_event.is_set():
|
||||
if processing_interrupted():
|
||||
print("Video generation interrupted by user")
|
||||
self._generation_interrupted = True
|
||||
|
||||
# Try to send interrupt signal to worker processes
|
||||
if self.generator is not None and hasattr(
|
||||
self.generator, 'executor'):
|
||||
try:
|
||||
# The MultiprocExecutor has a workers attribute
|
||||
if hasattr(self.generator.executor, 'workers'):
|
||||
for worker in self.generator.executor.workers:
|
||||
if worker.is_alive():
|
||||
os.kill(worker.pid, signal.SIGINT)
|
||||
print("Interrupt signal sent to worker processes")
|
||||
except Exception as e:
|
||||
print(f"Error sending interrupt signal: {e}")
|
||||
|
||||
# Set the interrupt event to notify other threads
|
||||
self._interrupt_event.set()
|
||||
break
|
||||
time.sleep(0.5)
|
||||
|
||||
def _run_generation(self, prompt: str, output_path: str,
|
||||
inference_args: dict[str, Any]) -> None:
|
||||
"""Thread function to run the generation"""
|
||||
try:
|
||||
if self.generator is not None:
|
||||
self.generator.generate_video(prompt=prompt,
|
||||
output_path=output_path,
|
||||
**inference_args)
|
||||
self._generation_result = os.path.join(output_path,
|
||||
f"{prompt[:100]}.mp4")
|
||||
else:
|
||||
raise RuntimeError("Generator is not initialized")
|
||||
except Exception as e:
|
||||
self._generation_exception = e
|
||||
self._interrupt_event.set()
|
||||
|
||||
def load_output_video(self, output_dir):
|
||||
video_extensions = ["*.mp4", "*.avi", "*.mov", "*.mkv"]
|
||||
video_files = []
|
||||
|
||||
for ext in video_extensions:
|
||||
video_files.extend(glob.glob(os.path.join(output_dir, ext)))
|
||||
|
||||
if not video_files:
|
||||
print("No video files found in output directory: %s", output_dir)
|
||||
return ""
|
||||
|
||||
video_files.sort()
|
||||
return video_files[0]
|
||||
|
||||
def launch_inference(
|
||||
self,
|
||||
prompt,
|
||||
output_path,
|
||||
num_gpus,
|
||||
model_path,
|
||||
embedded_cfg_scale,
|
||||
sp_size,
|
||||
tp_size,
|
||||
vae_precision,
|
||||
vae_tiling,
|
||||
vae_sp,
|
||||
text_encoder_precision,
|
||||
precision,
|
||||
inference_args=None,
|
||||
vae_config=None,
|
||||
text_encoder_config=None,
|
||||
dit_config=None,
|
||||
dit_cpu_offload=None,
|
||||
):
|
||||
print('Running FastVideo inference')
|
||||
|
||||
# Reset interruption flag and event
|
||||
self._generation_interrupted = False
|
||||
self._interrupt_event.clear()
|
||||
self._generation_result = None
|
||||
self._generation_exception = None
|
||||
|
||||
# Load pipeline config from model path
|
||||
pipeline_config = PipelineConfig.from_pretrained(model_path)
|
||||
print('pipeline_config', pipeline_config)
|
||||
|
||||
# Update configs with provided config dictionaries
|
||||
if dit_config is not None:
|
||||
update_config_from_args(pipeline_config.dit_config, dit_config)
|
||||
|
||||
if vae_config is not None:
|
||||
update_config_from_args(pipeline_config.vae_config, vae_config)
|
||||
|
||||
if text_encoder_config is not None:
|
||||
update_config_from_args(pipeline_config.text_encoder_configs,
|
||||
text_encoder_config)
|
||||
|
||||
# Update top-level pipeline config with remaining arguments
|
||||
raw_pipeline_args = {}
|
||||
if embedded_cfg_scale is not None:
|
||||
raw_pipeline_args['embedded_cfg_scale'] = embedded_cfg_scale
|
||||
if precision is not None:
|
||||
raw_pipeline_args['precision'] = precision
|
||||
if vae_precision is not None:
|
||||
raw_pipeline_args['vae_precision'] = vae_precision
|
||||
if vae_tiling is not None:
|
||||
raw_pipeline_args['vae_tiling'] = vae_tiling
|
||||
if vae_sp is not None:
|
||||
raw_pipeline_args['vae_sp'] = vae_sp
|
||||
if text_encoder_precision is not None:
|
||||
raw_pipeline_args['text_encoder_precision'] = text_encoder_precision
|
||||
|
||||
# Filter out any value explicitly set to -99999 (auto values)
|
||||
pipeline_args = {
|
||||
k: v
|
||||
for k, v in raw_pipeline_args.items() if str(int(v)) != str(-99999)
|
||||
}
|
||||
|
||||
update_config_from_args(pipeline_config, pipeline_args)
|
||||
|
||||
raw_generation_args = {}
|
||||
if num_gpus is not None:
|
||||
raw_generation_args['num_gpus'] = num_gpus
|
||||
if tp_size is not None:
|
||||
raw_generation_args['tp_size'] = tp_size
|
||||
if sp_size is not None:
|
||||
raw_generation_args['sp_size'] = sp_size
|
||||
if dit_cpu_offload is not None:
|
||||
raw_generation_args['dit_cpu_offload'] = dit_cpu_offload
|
||||
|
||||
generation_args = {
|
||||
k: v
|
||||
for k, v in raw_generation_args.items()
|
||||
if str(int(v)) != str(-99999)
|
||||
}
|
||||
|
||||
if self.generator is None:
|
||||
print('generation_args', generation_args)
|
||||
print('pipeline_config', pipeline_config)
|
||||
self.generator = FastVideoGenerator.from_pretrained(
|
||||
model_path=model_path,
|
||||
**generation_args,
|
||||
pipeline_config=pipeline_config)
|
||||
|
||||
print('inference_args', inference_args)
|
||||
|
||||
# Start a thread to run the generation
|
||||
self._generation_thread = threading.Thread(target=self._run_generation,
|
||||
args=(prompt, output_path,
|
||||
inference_args),
|
||||
daemon=True)
|
||||
self._generation_thread.start()
|
||||
|
||||
# Start a background thread to monitor for interruptions
|
||||
self._generation_active = True
|
||||
self._interrupt_thread = threading.Thread(
|
||||
target=self._monitor_for_interruption, daemon=True)
|
||||
self._interrupt_thread.start()
|
||||
|
||||
# Wait for either completion or interruption
|
||||
while self._generation_thread.is_alive(
|
||||
) and not self._interrupt_event.is_set():
|
||||
self._generation_thread.join(timeout=0.5)
|
||||
|
||||
self._generation_active = False
|
||||
if self._interrupt_thread:
|
||||
self._interrupt_thread.join(timeout=1.0)
|
||||
self._interrupt_thread = None
|
||||
|
||||
if self._generation_interrupted:
|
||||
print("Video generation was cancelled by user")
|
||||
raise GenerationCancelledException()
|
||||
elif self._generation_exception:
|
||||
# Re-raise the exception from the generation thread
|
||||
raise self._generation_exception
|
||||
elif self._generation_result:
|
||||
return (self._generation_result, )
|
||||
else:
|
||||
# This shouldn't happen, but just in case
|
||||
print("Generation completed but no result was produced")
|
||||
raise Exception("Generation failed to produce a result")
|
||||
@@ -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", "dit_cpu_offload", "enable_teacache"
|
||||
]
|
||||
const stringWidgetNames = ["prefix", "quant_config", "lora_config", "image_path"]
|
||||
|
||||
if (this.widgets) {
|
||||
for (let w of this.widgets) {
|
||||
if (intWidgetNames.includes(w.name)) {
|
||||
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOINT", { "default": 0 }]));
|
||||
} else if (floatWidgetNames.includes(w.name)) {
|
||||
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOFLOAT", { "default": 0 }]));
|
||||
} else if (comboWidgetNames.includes(w.name)) {
|
||||
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOCOMBO", { "default": 0 }]));
|
||||
} else if (stringWidgetNames.includes(w.name)) {
|
||||
new_widgets.push(makeAutoAnnotated(w, ["FVAUTOSTRING", { "default": "" }]));
|
||||
} else {
|
||||
new_widgets.push(w);
|
||||
}
|
||||
}
|
||||
this.widgets = new_widgets;
|
||||
|
||||
const autoWidgets = this.widgets.filter(w => w.type === "BOOLEAN" && w.isAuto !== undefined);
|
||||
}
|
||||
|
||||
this.graph?.setDirtyCanvas(true, true);
|
||||
});
|
||||
}
|
||||
},
|
||||
|
||||
async init() {
|
||||
// Force a redraw of all nodes when the extension initializes
|
||||
if (app.graph) {
|
||||
setTimeout(() => {
|
||||
app.graph.setDirtyCanvas(true, true);
|
||||
}, 1000);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
console.log("FastVideo.core.js loaded");
|
||||
@@ -1,11 +1,21 @@
|
||||
|
||||
|
||||
# Sliding Tile Atteniton Kernel
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
|
||||
## 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)
|
||||
|
||||
### Installation
|
||||
We support H100 (via TK) and any other GPU (via triton) for VSA.
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
@@ -15,27 +25,48 @@ sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
|
||||
## Environment Setup
|
||||
First, set up your CUDA environment:
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
git submodule update --init --recursive
|
||||
```
|
||||
|
||||
## Install Sliding Tile Attention (STA)
|
||||
### Verify if you have successfully installed
|
||||
|
||||
```bash
|
||||
# test numerical
|
||||
python tests/test_vsa.py
|
||||
# (For H100) test speed
|
||||
python benchmarks/bench_vsa_hopper.py
|
||||
```
|
||||
bench_vsa_hopper.py should print something like this:
|
||||
```bash
|
||||
Using topk=76 kv blocks per q block (out of 768 total kv blocks)
|
||||
|
||||
=== BLOCK SPARSE ATTENTION BENCHMARK ===
|
||||
Block Sparse Forward - TFLOPS: 5622.26
|
||||
Block Sparse Backward - TFLOPS: 3865.68
|
||||
```
|
||||
|
||||
|
||||
## Sliding Tile Attention (STA)
|
||||
We only support H100 for STA.
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup_sta.py install
|
||||
```
|
||||
|
||||
## Install Video Sparse Attention (VSA)
|
||||
|
||||
|
||||
|
||||
### Usage
|
||||
End-2-end inference with FastVideo:
|
||||
```bash
|
||||
python setup_vsa.py install
|
||||
bash scripts/inference/v1_inference_wan_STA.sh
|
||||
```
|
||||
|
||||
## Usage
|
||||
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.
|
||||
@@ -47,16 +78,21 @@ 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)
|
||||
|
||||
```
|
||||
|
||||
|
||||
## Test
|
||||
### Test
|
||||
```bash
|
||||
python test/test_sta.py
|
||||
python tests/test_sta.py # test STA
|
||||
python tests/test_vsa.py # test VSA
|
||||
```
|
||||
### Benchmark
|
||||
```bash
|
||||
python benchmarks/bench_sta.py
|
||||
```
|
||||
|
||||
## How Does STA Work?
|
||||
|
||||
### 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, [[3, 6, 10]] * 24, 0, False, '18x48x80')
|
||||
# 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, [[3, 6, 10]] * 24, 0, False, '18x48x80')
|
||||
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,12 @@ 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)
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize()
|
||||
# print(f"Average time for backward pass(ms): {ms:.2f}")
|
||||
# print(f"Average TFLOPS: {tflops_bwd}")
|
||||
# print("=" * 60)
|
||||
|
||||
return results
|
||||
|
||||
@@ -124,7 +115,10 @@ def plot_results(results):
|
||||
|
||||
# Example list of configurations to test
|
||||
configurations = [
|
||||
(2, 24, 69120, 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),
|
||||
@@ -1,9 +1,9 @@
|
||||
import torch
|
||||
import argparse
|
||||
from flash_attn.utils.benchmark import benchmark_forward
|
||||
from vsa import block_sparse_attention_fwd, block_sparse_attention_backward
|
||||
from triton.testing import do_bench
|
||||
from vsa import block_sparse_fwd, block_sparse_bwd
|
||||
from vsa import BLOCK_M, BLOCK_N
|
||||
|
||||
import triton
|
||||
import numpy as np
|
||||
import random
|
||||
|
||||
@@ -23,7 +23,7 @@ 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('--head_dim', type=int, default=128, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=None, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[49152], help='Sequence lengths to benchmark')
|
||||
return parser.parse_args()
|
||||
@@ -130,19 +130,19 @@ def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_
|
||||
|
||||
# Forward pass
|
||||
# Warm-up run
|
||||
o, l_vec = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
variable_block_sizes = torch.ones(q2k_block_sparse_index.shape[2], device=q.device).int() * BLOCK_M
|
||||
o, l_vec = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Benchmark forward
|
||||
_, fwd_time = benchmark_forward(
|
||||
block_sparse_attention_fwd,
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num,
|
||||
repeats=20,
|
||||
verbose=False,
|
||||
desc='Block Sparse Forward'
|
||||
fwd_time = do_bench(
|
||||
lambda: block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes),
|
||||
warmup=5,
|
||||
rep=20,
|
||||
quantiles=None
|
||||
)
|
||||
|
||||
sparse_tflops = flops / fwd_time.mean * 1e-12
|
||||
sparse_tflops = flops / fwd_time * 1e-12 * 1e3
|
||||
print(f"Block Sparse Forward - TFLOPS: {sparse_tflops:.2f}")
|
||||
|
||||
# Backward pass
|
||||
@@ -150,20 +150,19 @@ def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_
|
||||
|
||||
# Warm-up runs
|
||||
for _ in range(5):
|
||||
block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Benchmark backward
|
||||
_, bwd_time = benchmark_forward(
|
||||
block_sparse_attention_backward,
|
||||
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_time = do_bench(
|
||||
lambda: block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes),
|
||||
warmup=5,
|
||||
rep=20,
|
||||
quantiles=None
|
||||
)
|
||||
bwd_flops = 2.5 * flops # Approximation
|
||||
|
||||
sparse_bwd_tflops = bwd_flops / bwd_time.mean * 1e-12
|
||||
sparse_bwd_tflops = bwd_flops / bwd_time * 1e-12 * 1e3
|
||||
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd_tflops:.2f}")
|
||||
|
||||
return sparse_tflops, sparse_bwd_tflops
|
||||
@@ -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,2 +1,2 @@
|
||||
recursive-include tk *
|
||||
include config.py
|
||||
include config_sta.py
|
||||
@@ -0,0 +1,87 @@
|
||||
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Sliding Tile Attention (STA)
|
||||
We only support H100 for STA.
|
||||
|
||||
### Installation
|
||||
```bash
|
||||
pip install st_attn
|
||||
```
|
||||
|
||||
Install from source:
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
### Usage
|
||||
End-2-end inference with FastVideo:
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_STA.sh
|
||||
```
|
||||
|
||||
If you want to use sliding tile attention in your custom model:
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
# q, k, v: [batch_size, num_heads, seq_length, head_dim], seq_length = T*H*W + 256
|
||||
# a tile is a cube of size (6, 8, 8)
|
||||
# window_size in tiles: [(window_t, window_h, window_w), (..)...]. For example, window size (3, 3, 3) means a query can attend to (3x6, 3x8, 3x8) = (18, 24, 24) tokens out of the total 30x48x80 video.
|
||||
# text_length: int ranging from 0 to 256
|
||||
# If your attention contains text token (Hunyuan)
|
||||
out = sliding_tile_attention(q, k, v, window_size, text_length)
|
||||
# If your attention does not contain text token (StepVideo)
|
||||
out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
```
|
||||
|
||||
|
||||
### Test
|
||||
```bash
|
||||
python ../tests/test_sta.py # test STA
|
||||
python ../tests/test_vsa.py # test VSA
|
||||
```
|
||||
### Benchmark
|
||||
```bash
|
||||
python ../benchmarks/bench_sta.py
|
||||
```
|
||||
|
||||
|
||||
### How Does STA Work?
|
||||
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
|
||||
https://github.com/user-attachments/assets/f3b6dd79-7b43-4b60-a0fa-3d6495ec5747
|
||||
|
||||
## Why is STA Fast?
|
||||
2D/3D Sliding Window Attention (SWA) creates many mixed blocks in the attention map. Even though mixed blocks have less output value,a mixed block is significantly slower than a dense block due to the GPU-unfriendly masking operation.
|
||||
|
||||
STA removes mixed blocks.
|
||||
|
||||
|
||||
<div align="center">
|
||||
<img src=../../../assets/sliding_tile_attn_map.png width="80%"/>
|
||||
</div>
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
from csrc.attn.config_sta import kernels, sources, target
|
||||
from config_sta import kernels, sources, target
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
@@ -9,7 +9,7 @@ target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "st_attn"
|
||||
VERSION = "0.0.4"
|
||||
VERSION = "0.0.6"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Sliding Tile Atteniton Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/sliding_tile_attention"
|
||||
@@ -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,10 +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();
|
||||
}
|
||||
|
||||
@@ -1,266 +0,0 @@
|
||||
import torch
|
||||
import argparse
|
||||
from flash_attn.utils.benchmark import benchmark_forward
|
||||
from flash_attn import flash_attn_func
|
||||
from vsa import block_sparse_attention_fwd, block_sparse_attention_backward, BlockSparseAttentionFunction
|
||||
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=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=100, help='Number of test iterations to run')
|
||||
return parser.parse_args()
|
||||
|
||||
@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 = parse_arguments()
|
||||
|
||||
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_sdpa = q.clone()
|
||||
k_sdpa = k.clone()
|
||||
v_sdpa = v.clone()
|
||||
|
||||
q.requires_grad = True
|
||||
k.requires_grad = True
|
||||
v.requires_grad = True
|
||||
q_sdpa.requires_grad = True
|
||||
k_sdpa.requires_grad = True
|
||||
v_sdpa.requires_grad = True
|
||||
|
||||
# testing forward
|
||||
o = BlockSparseAttentionFunction.apply(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
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)
|
||||
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
|
||||
grad_o = torch.randn_like(o)
|
||||
o.backward(grad_o)
|
||||
o_sdpa.backward(grad_o)
|
||||
|
||||
sim, l1, rmse = precision_metric(q.grad, q_sdpa.grad)
|
||||
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)
|
||||
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)
|
||||
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}")
|
||||
|
||||
# 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}")
|
||||
print(f"L1 error: mean={np.mean(forward_metrics['l1']):.6f}, std={np.std(forward_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(forward_metrics['rmse']):.6f}, std={np.std(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}")
|
||||
print(f"L1 error: mean={np.mean(grad_q_metrics['l1']):.6f}, std={np.std(grad_q_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_q_metrics['rmse']):.6f}, std={np.std(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}")
|
||||
print(f"L1 error: mean={np.mean(grad_k_metrics['l1']):.6f}, std={np.std(grad_k_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_k_metrics['rmse']):.6f}, std={np.std(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}")
|
||||
print(f"L1 error: mean={np.mean(grad_v_metrics['l1']):.6f}, std={np.std(grad_v_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_v_metrics['rmse']):.6f}, std={np.std(grad_v_metrics['rmse']):.6f}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,136 +0,0 @@
|
||||
import torch
|
||||
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 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, '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':
|
||||
tensors_fa2_pt = [(pt_o, fa2_o)]
|
||||
else: # 'forward_backward'
|
||||
if error_mode == 'output':
|
||||
tensors_fa2_pt = [(pt_o, fa2_o)]
|
||||
elif error_mode == 'backward':
|
||||
tensors_fa2_pt = [(pt_qg, fa2_qg),
|
||||
(pt_kg, fa2_kg),
|
||||
(pt_vg, fa2_vg)]
|
||||
else: # 'all'
|
||||
tensors_fa2_pt = [(pt_o, fa2_o),
|
||||
(pt_qg, fa2_qg),
|
||||
(pt_kg, fa2_kg),
|
||||
(pt_vg, fa2_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())
|
||||
|
||||
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"ATTENTION 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} | {'FA2 vs PT Avg':<15} | {'FA2 vs PT Max':<15}")
|
||||
print(f"{'-'*12} | {'-'*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']
|
||||
|
||||
# Print row
|
||||
print(f"{n:<12} | {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("Attention error comparison completed.")
|
||||
@@ -1,175 +0,0 @@
|
||||
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.")
|
||||
@@ -81,5 +81,7 @@ std = 10
|
||||
|
||||
# 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,156 @@
|
||||
import torch
|
||||
import sys
|
||||
import os
|
||||
import numpy as np
|
||||
from tqdm import tqdm
|
||||
|
||||
# Add the parent directory to the path to import block_sparse_attn
|
||||
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
from tests.utils import generate_block_sparse_mask_for_function, create_full_mask_from_block_mask
|
||||
from vsa import block_sparse_attn
|
||||
|
||||
BLOCK_M = 64
|
||||
BLOCK_N = 64
|
||||
|
||||
def pytorch_test(Q, K, V, block_sparse_mask, dO):
|
||||
q_ = Q.clone().float().requires_grad_()
|
||||
k_ = K.clone().float().requires_grad_()
|
||||
v_ = V.clone().float().requires_grad_()
|
||||
|
||||
QK = torch.matmul(q_, k_.transpose(-2, -1))
|
||||
QK /= (q_.size(-1) ** 0.5)
|
||||
QK = QK.masked_fill(~block_sparse_mask.unsqueeze(0), float('-inf'))
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v_)
|
||||
|
||||
dO_ = dO
|
||||
output.backward(dO_)
|
||||
return (
|
||||
output.to(torch.bfloat16),
|
||||
q_.grad.to(torch.bfloat16),
|
||||
k_.grad.to(torch.bfloat16),
|
||||
v_.grad.to(torch.bfloat16),
|
||||
)
|
||||
|
||||
|
||||
def block_sparse_kernel_test(Q, K, V, block_sparse_mask, variable_block_sizes, non_pad_index, dO):
|
||||
Q = Q.detach().requires_grad_()
|
||||
K = K.detach().requires_grad_()
|
||||
V = V.detach().requires_grad_()
|
||||
|
||||
q_padded = vsa_pad(Q, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
|
||||
k_padded = vsa_pad(K, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
|
||||
v_padded = vsa_pad(V, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
|
||||
output, _= block_sparse_attn(q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes)
|
||||
output = output[:, :, non_pad_index, :]
|
||||
output.backward(dO)
|
||||
return output, Q.grad, K.grad, V.grad
|
||||
|
||||
|
||||
def get_non_pad_index(
|
||||
vid_len: torch.LongTensor,
|
||||
n_win: int,
|
||||
win_size: int,
|
||||
):
|
||||
device = vid_len.device
|
||||
starts_pad = torch.arange(n_win, device=device) * win_size
|
||||
index_pad = starts_pad[:, None] + torch.arange(win_size, device=device)[None, :]
|
||||
index_mask = torch.arange(win_size, device=device)[None, :] < vid_len[:, None]
|
||||
|
||||
return index_pad[index_mask]
|
||||
|
||||
def generate_tensor(shape, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
return tensor
|
||||
|
||||
def generate_variable_block_sizes(num_blocks, min_size=32, max_size=64, device="cuda"):
|
||||
return torch.randint(min_size, max_size + 1, (num_blocks,), device=device, dtype=torch.int32)
|
||||
|
||||
|
||||
def vsa_pad(x, non_pad_index, num_blocks, block_size):
|
||||
padded_x = torch.zeros((1, x.shape[1], num_blocks * BLOCK_M, x.shape[3]), device=x.device, dtype=x.dtype)
|
||||
padded_x[:, :, non_pad_index, :] = x
|
||||
return padded_x
|
||||
|
||||
def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all'):
|
||||
results = {
|
||||
'gO': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
'gQ': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
'gK': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
'gV': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
|
||||
}
|
||||
|
||||
device = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
variable_block_sizes = generate_variable_block_sizes(num_blocks, device=device)
|
||||
S = int(variable_block_sizes.sum().item())
|
||||
padded_S = num_blocks * BLOCK_M
|
||||
non_pad_index = get_non_pad_index(variable_block_sizes, num_blocks, BLOCK_M)
|
||||
block_mask = generate_block_sparse_mask_for_function(h, num_blocks, k, device)
|
||||
full_mask = create_full_mask_from_block_mask(block_mask, variable_block_sizes, device)
|
||||
for _ in range(num_iterations):
|
||||
Q = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
K = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
V = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
dO = generate_tensor((1, h, S, d), torch.bfloat16, device)
|
||||
|
||||
# dO_padded = torch.zeros_like(dO_padded)
|
||||
# dO_padded[:, :, non_pad_index, :] = dO
|
||||
|
||||
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, full_mask, dO)
|
||||
bs_o, bs_qg, bs_kg, bs_vg = block_sparse_kernel_test(Q, K, V, block_mask.unsqueeze(0), variable_block_sizes,non_pad_index, dO)
|
||||
for name, (pt, bs) in zip(['gQ', 'gK', 'gV', 'gO'], [(pt_qg, bs_qg), (pt_kg, bs_kg), (pt_vg, bs_vg), (pt_o, bs_o)]):
|
||||
if bs is not None:
|
||||
diff = pt - bs
|
||||
abs_diff = torch.abs(diff)
|
||||
results[name]['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results[name]['sum_abs'] += torch.sum(torch.abs(pt)).item()
|
||||
rel_max_diff = torch.max(abs_diff) / torch.mean(torch.abs(pt))
|
||||
results[name]['max_diff'] = max(results[name]['max_diff'], rel_max_diff.item())
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
total_elements = h * S * d * num_iterations
|
||||
for name, data in results.items():
|
||||
avg_diff = data['sum_diff'] / total_elements
|
||||
max_diff = data['max_diff']
|
||||
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
|
||||
|
||||
return results
|
||||
|
||||
def generate_error_graphs(h, d, error_mode='all'):
|
||||
test_configs = [
|
||||
{"num_blocks": 16, "k": 2, "description": "Small sequence"},
|
||||
{"num_blocks": 32, "k": 4, "description": "Medium sequence"},
|
||||
{"num_blocks": 53, "k": 6, "description": "Large sequence"},
|
||||
]
|
||||
|
||||
print(f"\nError Analysis for h={h}, d={d}, mode={error_mode}")
|
||||
print("=" * 150)
|
||||
print(f"{'Config':<20} {'Blocks':<8} {'K':<4} "
|
||||
f"{'gQ Avg':<12} {'Rel gQ Max':<12} "
|
||||
f"{'gK Avg':<12} {'Rel gK Max':<12} "
|
||||
f"{'gV Avg':<12} {'Rel gV Max':<12} "
|
||||
f"{'gO Avg':<12} {'Rel gO Max':<12}")
|
||||
print("-" * 150)
|
||||
|
||||
for config in test_configs:
|
||||
num_blocks = config["num_blocks"]
|
||||
k = config["k"]
|
||||
description = config["description"]
|
||||
results = check_correctness(h, d, num_blocks, k, error_mode=error_mode)
|
||||
print(f"{description:<20} {num_blocks:<8} {k:<4} "
|
||||
f"{results['gQ']['avg_diff']:<12.6e} {results['gQ']['max_diff']:<12.6e} "
|
||||
f"{results['gK']['avg_diff']:<12.6e} {results['gK']['max_diff']:<12.6e} "
|
||||
f"{results['gV']['avg_diff']:<12.6e} {results['gV']['max_diff']:<12.6e} "
|
||||
f"{results['gO']['avg_diff']:<12.6e} {results['gO']['max_diff']:<12.6e}")
|
||||
|
||||
print("-" * 150)
|
||||
|
||||
if __name__ == "__main__":
|
||||
h, d = 16, 128
|
||||
print("Block Sparse Attention with Variable Block Sizes Analysis")
|
||||
print("=" * 60)
|
||||
for mode in ['backward']:
|
||||
generate_error_graphs(h, d, error_mode=mode)
|
||||
print("\nAnalysis completed for all modes.")
|
||||
@@ -0,0 +1,54 @@
|
||||
import torch
|
||||
|
||||
def generate_block_sparse_mask_for_function(h, num_blocks, k, device="cuda"):
|
||||
"""
|
||||
Generate block sparse mask of shape [h, num_blocks, num_blocks].
|
||||
|
||||
Args:
|
||||
h: number of heads
|
||||
num_blocks: number of blocks
|
||||
k: number of kv blocks each q block attends to
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
|
||||
"""
|
||||
k = min(k, num_blocks)
|
||||
scores = torch.rand(h, num_blocks, num_blocks, device=device)
|
||||
_, indices = torch.topk(scores, k, dim=-1)
|
||||
block_sparse_mask = torch.zeros(h, num_blocks, num_blocks, dtype=torch.bool, device=device)
|
||||
|
||||
block_sparse_mask = block_sparse_mask.scatter_(2, indices, 1).bool()
|
||||
return block_sparse_mask
|
||||
|
||||
|
||||
def create_full_mask_from_block_mask(block_sparse_mask, variable_block_sizes, device="cuda"):
|
||||
"""
|
||||
Convert block-level sparse mask to full attention mask.
|
||||
|
||||
Args:
|
||||
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
|
||||
variable_block_sizes: [num_blocks] tensor
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
full_mask: [h, S, S] bool tensor where S = total sequence length
|
||||
"""
|
||||
h, num_blocks, _ = block_sparse_mask.shape
|
||||
total_seq_len = variable_block_sizes.sum().item()
|
||||
cumsum = torch.cat([torch.tensor([0], device=device), variable_block_sizes.cumsum(dim=0)[:-1]])
|
||||
|
||||
full_mask = torch.zeros(h, total_seq_len, total_seq_len, dtype=torch.bool, device=device)
|
||||
|
||||
for head in range(h):
|
||||
for q_block in range(num_blocks):
|
||||
q_start = cumsum[q_block]
|
||||
q_end = q_start + variable_block_sizes[q_block]
|
||||
|
||||
for kv_block in range(num_blocks):
|
||||
if block_sparse_mask[head, q_block, kv_block]:
|
||||
kv_start = cumsum[kv_block]
|
||||
kv_end = kv_start + variable_block_sizes[kv_block]
|
||||
full_mask[head, q_start:q_end, kv_start:kv_end] = True
|
||||
|
||||
return full_mask
|
||||
@@ -0,0 +1,2 @@
|
||||
recursive-include tk *
|
||||
include config_vsa.py
|
||||
@@ -0,0 +1,61 @@
|
||||
|
||||
|
||||
# Attention Kernel Used in FastVideo
|
||||
|
||||
## Video Sparse Attention (VSA)
|
||||
|
||||
### Installation
|
||||
We support H100 (via TK) and any other GPU (via triton) for VSA.
|
||||
|
||||
```bash
|
||||
pip install vsa
|
||||
```
|
||||
|
||||
Install from source:
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
|
||||
If you encounter error during installation, try below:
|
||||
Install C++20 for ThunderKittens:
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
(If you use CUDA12.8)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.8
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
```
|
||||
|
||||
### Verify if you have successfully installed
|
||||
|
||||
```bash
|
||||
# test numerical
|
||||
python ../tests/test_vsa.py
|
||||
# (For H100) test speed
|
||||
python ../benchmarks/bench_vsa_hopper.py
|
||||
```
|
||||
|
||||
bench_vsa_hopper.py should print something like this:
|
||||
|
||||
```bash
|
||||
Using topk=76 kv blocks per q block (out of 768 total kv blocks)
|
||||
|
||||
=== BLOCK SPARSE ATTENTION BENCHMARK ===
|
||||
Block Sparse Forward - TFLOPS: 5622.26
|
||||
Block Sparse Backward - TFLOPS: 3865.68
|
||||
```
|
||||
|
||||
## Acknowledgement
|
||||
|
||||
We learned or reuse code from FlexAtteniton, NATEN, and ThunderKittens.
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
from csrc.attn.config_vsa import kernels, sources, target
|
||||
from config_vsa import kernels, sources, target
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
@@ -9,10 +9,10 @@ target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "vsa"
|
||||
VERSION = "0.0.1"
|
||||
VERSION = "0.0.3"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Video Sparse Attention Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn/video_sparse_attn"
|
||||
|
||||
# Set environment variables
|
||||
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
|
||||
@@ -51,21 +51,26 @@ for k in kernels:
|
||||
source_files.append(sources[k]['source_files'][target])
|
||||
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
|
||||
|
||||
|
||||
ext_modules = [
|
||||
CUDAExtension('vsa_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
]
|
||||
|
||||
|
||||
|
||||
setup(name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
ext_modules=[
|
||||
CUDAExtension('vsa_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
],
|
||||
ext_modules=ext_modules,
|
||||
cmdclass={'build_ext': BuildExtension},
|
||||
classifiers=[
|
||||
"Programming Language :: Python :: 3",
|
||||
@@ -9,10 +9,10 @@
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE
|
||||
extern std::vector<torch::Tensor> block_sparse_attention_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num, torch::Tensor block_size
|
||||
);
|
||||
extern std::vector<torch::Tensor> block_sparse_attention_backward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, torch::Tensor l_vec, torch::Tensor og, torch::Tensor k2q_block_sparse_index, torch::Tensor k2q_block_sparse_num
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, torch::Tensor l_vec, torch::Tensor og, torch::Tensor k2q_block_sparse_index, torch::Tensor k2q_block_sparse_num, torch::Tensor block_size
|
||||
);
|
||||
#endif
|
||||
|
||||
@@ -0,0 +1,80 @@
|
||||
import torch
|
||||
from typing import Tuple
|
||||
block_sparse_attn=None
|
||||
import torch
|
||||
major, minor = torch.cuda.get_device_capability(0)
|
||||
if major == 9 and minor == 0:# check if H100
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
from vsa.block_sparse_wrapper import block_sparse_attn_SM90
|
||||
block_sparse_attn = block_sparse_attn_SM90
|
||||
else:
|
||||
from vsa.block_sparse_wrapper import block_sparse_attn_triton
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
block_sparse_attn = block_sparse_attn_triton
|
||||
|
||||
BLOCK_M = 64
|
||||
BLOCK_N = 64
|
||||
|
||||
|
||||
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
QK = torch.matmul(q, k.transpose(-2, -1))
|
||||
QK /= (q.size(-1)**0.5)
|
||||
|
||||
# Causal mask removed since causal is always false
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v)
|
||||
return output, QK
|
||||
|
||||
|
||||
def video_sparse_attn(q, k, v, variable_block_sizes, topk, block_size, compress_attn_weight=None):
|
||||
"""
|
||||
q: [batch_size, num_heads, seq_len, head_dim]
|
||||
k: [batch_size, num_heads, seq_len, head_dim]
|
||||
v: [batch_size, num_heads, seq_len, head_dim]
|
||||
topk: int
|
||||
block_size: int or tuple of 3 ints
|
||||
video_shape: tuple of (T, H, W)
|
||||
compress_attn_weight: [batch_size, num_heads, seq_len, head_dim]
|
||||
select_attn_weight: [batch_size, num_heads, seq_len, head_dim]
|
||||
NOTE: We assume q, k, v is zero padded!!
|
||||
V1 of sparse attention. Include compress attn and sparse attn branch, use average pooling to compress.
|
||||
Assume q, k, v is flattened in this way: [batch_size, num_heads, T//block_size[0], H//block_size[1], W//block_size[2], block_size[0], block_size[1], block_size[2]]
|
||||
"""
|
||||
|
||||
if isinstance(block_size, int):
|
||||
block_size = (block_size, block_size, block_size)
|
||||
|
||||
block_elements = block_size[0] * block_size[1] * block_size[2]
|
||||
assert block_elements == 64
|
||||
assert q.shape[2] % block_elements == 0
|
||||
batch_size, num_heads, seq_len, head_dim = q.shape
|
||||
# compress attn
|
||||
q_compress = (q.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(q.dtype)
|
||||
k_compress = (k.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(k.dtype)
|
||||
v_compress = (v.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).float().sum(dim=3) / variable_block_sizes.view(1, 1, -1, 1)).to(v.dtype)
|
||||
|
||||
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
|
||||
v_compress)
|
||||
|
||||
output_compress = output_compress.view(batch_size, num_heads,
|
||||
seq_len // block_elements, 1,
|
||||
head_dim)
|
||||
output_compress = output_compress.repeat(1, 1, 1, block_elements,
|
||||
1).view(batch_size, num_heads,
|
||||
seq_len, head_dim)
|
||||
|
||||
topK_indices = torch.topk(block_attn_score, topk, dim=-1).indices
|
||||
block_mask = torch.zeros_like(block_attn_score, dtype=torch.bool).scatter_(-1, topK_indices, True)
|
||||
output_select, _ = block_sparse_attn(q, k, v, block_mask, variable_block_sizes)
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
final_output = output_compress * compress_attn_weight + output_select
|
||||
else:
|
||||
final_output = output_compress + output_select
|
||||
return final_output
|
||||
|
||||
@@ -0,0 +1,449 @@
|
||||
"""
|
||||
Fused Attention
|
||||
===============
|
||||
|
||||
This is a Triton implementation of the Flash Attention v2 algorithm from Tri Dao
|
||||
(https://tridao.me/publications/flash2/flash2.pdf)
|
||||
|
||||
Credits: OpenAI kernel team
|
||||
"""
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
import math # small utility needed by the sparse wrapper
|
||||
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
|
||||
|
||||
|
||||
# We don't run auto-tuning every time to keep the tutorial fast. Keeping
|
||||
# the code below and commenting out the equivalent parameters is convenient for
|
||||
# re-tuning.
|
||||
configs = [
|
||||
triton.Config({'BLOCK_M': BM, 'BLOCK_N': BN}, num_stages=s, num_warps=w) \
|
||||
for BM in [64]\
|
||||
for BN in [64]\
|
||||
for s in [3, 4, 7]\
|
||||
for w in [4, 8]\
|
||||
]
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
@triton.autotune(configs, key=["N_CTX", "HEAD_DIM"])
|
||||
@triton.jit
|
||||
def _attn_fwd_sparse(Q, K, V, sm_scale, #
|
||||
q2k_index, q2k_num, max_kv_blks, #
|
||||
variable_block_sizes,
|
||||
M, Out, #
|
||||
stride_qz, stride_qh, stride_qm, stride_qk,
|
||||
stride_kz, stride_kh, stride_kn, stride_kk,
|
||||
stride_vz, stride_vh, stride_vk, stride_vn,
|
||||
stride_oz, stride_oh, stride_om, stride_on,
|
||||
Z, H, N_CTX, #
|
||||
HEAD_DIM: tl.constexpr, #
|
||||
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
|
||||
STAGE: tl.constexpr):
|
||||
"""
|
||||
64×64 **block-sparse** forward kernel. Back-prop kernels remain dense
|
||||
(32×64 and 64×32) – memory footprint unchanged.
|
||||
"""
|
||||
|
||||
# ----- program-id mapping -----
|
||||
q_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(1) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_M
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
|
||||
# ----- base pointers -----
|
||||
qvk_off = (b.to(tl.int64) * stride_qz +
|
||||
h.to(tl.int64) * stride_qh)
|
||||
|
||||
Q_ptr = tl.make_block_ptr(
|
||||
base=Q + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_qm, stride_qk),
|
||||
offsets=(q_blk * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
|
||||
|
||||
K_base = tl.make_block_ptr(
|
||||
base=K + qvk_off, shape=(HEAD_DIM, N_CTX),
|
||||
strides=(stride_kk, stride_kn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(HEAD_DIM, BLOCK_N), order=(0, 1))
|
||||
|
||||
v_order: tl.constexpr = (0, 1) if V.dtype.element_ty == tl.float8e5 else (1, 0)
|
||||
V_base = tl.make_block_ptr(
|
||||
base=V + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_vk, stride_vn),
|
||||
offsets=(0, 0),
|
||||
block_shape=(BLOCK_N, HEAD_DIM), order=v_order)
|
||||
|
||||
O_ptr = tl.make_block_ptr(
|
||||
base=Out + qvk_off, shape=(N_CTX, HEAD_DIM),
|
||||
strides=(stride_om, stride_on),
|
||||
offsets=(q_blk * BLOCK_M, 0),
|
||||
block_shape=(BLOCK_M, HEAD_DIM), order=(1, 0))
|
||||
|
||||
# ----- accumulators -----
|
||||
offs_m = q_blk * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
m_i = tl.full([BLOCK_M], -float("inf"), tl.float32)
|
||||
l_i = tl.zeros([BLOCK_M], dtype=tl.float32) + 1.0
|
||||
acc = tl.zeros([BLOCK_M, HEAD_DIM], dtype=tl.float32)
|
||||
qk_scale = sm_scale * 1.44269504 # 1/ln2
|
||||
q = tl.load(Q_ptr)
|
||||
|
||||
# ----- sparse loop over valid K/V tiles -----
|
||||
for i in range(0, kv_blocks):
|
||||
kv_idx = tl.load(kv_ptr + i).to(tl.int32)
|
||||
block_size = tl.load(variable_block_sizes + kv_idx)
|
||||
K_ptr = tl.advance(K_base, (0, kv_idx * BLOCK_N))
|
||||
V_ptr = tl.advance(V_base, (kv_idx * BLOCK_N, 0))
|
||||
|
||||
k = tl.load(K_ptr)
|
||||
qk = tl.dot(q, k)
|
||||
# mask out invalid columns
|
||||
mask = tl.arange(0, BLOCK_N) < block_size
|
||||
qk = tl.where(mask[None, :], qk, -float("inf"))
|
||||
|
||||
m_ij = tl.maximum(m_i, tl.max(qk, 1) * qk_scale)
|
||||
p = tl.math.exp2(qk * qk_scale - m_ij[:, None])
|
||||
l_ij = tl.sum(p, 1)
|
||||
|
||||
alpha = tl.math.exp2(m_i - m_ij)
|
||||
l_i = l_i * alpha + l_ij
|
||||
acc = acc * alpha[:, None]
|
||||
|
||||
v = tl.load(V_ptr)
|
||||
acc = tl.dot(p.to(tl.bfloat16), v, acc)
|
||||
m_i = m_ij
|
||||
|
||||
# ----- epilogue -----
|
||||
m_i += tl.math.log2(l_i)
|
||||
acc = acc / l_i[:, None]
|
||||
tl.store(M + off_hz * N_CTX + offs_m, m_i)
|
||||
tl.store(O_ptr, acc.to(Out.type.element_ty))
|
||||
# ──────────────────────────── SPARSE ADDITION END ─────────────────────────────
|
||||
|
||||
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd_preprocess(O, DO, #
|
||||
Delta, #
|
||||
Z, H, N_CTX, #
|
||||
BLOCK_M: tl.constexpr, HEAD_DIM: tl.constexpr #
|
||||
):
|
||||
off_m = tl.program_id(0) * BLOCK_M + tl.arange(0, BLOCK_M)
|
||||
off_hz = tl.program_id(1)
|
||||
off_n = tl.arange(0, HEAD_DIM)
|
||||
# load
|
||||
o = tl.load(O + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :])
|
||||
do = tl.load(DO + off_hz * HEAD_DIM * N_CTX + off_m[:, None] * HEAD_DIM + off_n[None, :]).to(tl.float32)
|
||||
delta = tl.sum(o * do, axis=1)
|
||||
# write-back
|
||||
tl.store(Delta + off_hz * N_CTX + off_m, delta)
|
||||
|
||||
|
||||
# The main inner-loop logic for computing dK and dV.
|
||||
@triton.jit
|
||||
def _attn_bwd_dkdv(dk, dv, #
|
||||
Q, k, v, sm_scale, #
|
||||
DO, #
|
||||
M, D, #
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
# shared by Q/K/V/DO.
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, BLOCK_M1: tl.constexpr, #
|
||||
BLOCK_N1: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr, #
|
||||
# Filled in by the wrapper.
|
||||
start_n, start_m, num_steps):
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M1)
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N1)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
qT_ptrs = Q + offs_m[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
do_ptrs = DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
# BLOCK_N1 must be a multiple of BLOCK_M1, otherwise the code wouldn't work.
|
||||
tl.static_assert(BLOCK_N1 % BLOCK_M1 == 0)
|
||||
step_m = BLOCK_M1
|
||||
kv_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(2) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_N1
|
||||
meta_base = ((b * H + h) * q_tiles + kv_blk)
|
||||
|
||||
q_blocks = tl.load(k2q_num + meta_base) # int32
|
||||
q_ptr = k2q_index + meta_base * max_q_blks # ptr to list
|
||||
block_size = tl.load(variable_block_sizes + kv_blk)
|
||||
|
||||
|
||||
|
||||
for blk_idx in range(q_blocks*2):
|
||||
block_sparse_offset = (tl.load(q_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_m
|
||||
qT = tl.load(qT_ptrs + block_sparse_offset * stride_tok)
|
||||
# Load m before computing qk to reduce pipeline stall.
|
||||
offs_m = start_m + block_sparse_offset + tl.arange(0, BLOCK_M1)
|
||||
m = tl.load(M + offs_m)
|
||||
qkT = tl.dot(k, qT)
|
||||
pT = tl.math.exp2(qkT - m[None, :])
|
||||
mask = tl.arange(0, BLOCK_N1) < block_size
|
||||
pT = tl.where(mask[:, None], pT, 0.0)
|
||||
|
||||
do = tl.load(do_ptrs + block_sparse_offset * stride_tok)
|
||||
# Compute dV.
|
||||
ppT = pT
|
||||
ppT = ppT.to(tl.bfloat16)
|
||||
dv += tl.dot(ppT, do)
|
||||
# D (= delta) is pre-divided by ds_scale.
|
||||
Di = tl.load(D + offs_m)
|
||||
# Compute dP and dS.
|
||||
dpT = tl.dot(v, tl.trans(do)).to(tl.float32)
|
||||
dsT = pT * (dpT - Di[None, :])
|
||||
dsT = dsT.to(tl.bfloat16)
|
||||
dk += tl.dot(dsT, tl.trans(qT))
|
||||
# Increment pointers.
|
||||
return dk, dv
|
||||
|
||||
|
||||
|
||||
# the main inner-loop logic for computing dQ
|
||||
@triton.jit
|
||||
def _attn_bwd_dq(dq, q, K, V, #
|
||||
do, m, D,
|
||||
# shared by Q/K/V/DO.
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M2: tl.constexpr, #
|
||||
BLOCK_N2: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr,
|
||||
# Filled in by the wrapper.
|
||||
start_m, start_n, num_steps):
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M2)
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N2)
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
kT_ptrs = K + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
vT_ptrs = V + offs_n[None, :] * stride_tok + offs_k[:, None] * stride_d
|
||||
# D (= delta) is pre-divided by ds_scale.
|
||||
Di = tl.load(D + offs_m)
|
||||
# BLOCK_M2 must be a multiple of BLOCK_N2, otherwise the code wouldn't work.
|
||||
tl.static_assert(BLOCK_M2 % BLOCK_N2 == 0)
|
||||
step_n = BLOCK_N2
|
||||
|
||||
q_blk = tl.program_id(0) # Q-tile index
|
||||
off_hz = tl.program_id(2) # fused (batch, head)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
q_tiles = N_CTX // BLOCK_M2
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
|
||||
kv_blocks = tl.load(q2k_num + meta_base) # int32
|
||||
kv_ptr = q2k_index + meta_base * max_kv_blks # ptr to list
|
||||
|
||||
|
||||
for blk_idx in range(kv_blocks*2):
|
||||
block_sparse_offset = (tl.load(kv_ptr + blk_idx//2).to(tl.int32)*2 + blk_idx%2) *step_n * stride_tok
|
||||
block_size = tl.load(variable_block_sizes + blk_idx//2) - (blk_idx%2) * step_n
|
||||
kT = tl.load(kT_ptrs + block_sparse_offset)
|
||||
vT = tl.load(vT_ptrs + block_sparse_offset)
|
||||
qk = tl.dot(q, kT)
|
||||
p = tl.math.exp2(qk - m)
|
||||
mask = tl.arange(0, BLOCK_N2) < block_size.to(tl.int32)
|
||||
p = tl.where(mask[None, :], p , 0.0)
|
||||
# Compute dP and dS.
|
||||
dp = tl.dot(do, vT).to(tl.float32)
|
||||
ds = p * (dp - Di[:, None])
|
||||
ds = ds.to(tl.bfloat16)
|
||||
# Compute dQ.
|
||||
# NOTE: We need to de-scale dq in the end, because kT was pre-scaled.
|
||||
dq += tl.dot(ds, tl.trans(kT))
|
||||
# Increment pointers.
|
||||
return dq
|
||||
|
||||
|
||||
|
||||
@triton.jit
|
||||
def _attn_bwd(Q, K, V, sm_scale, #
|
||||
DO, #
|
||||
DQ, DK, DV, #
|
||||
M, D,
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
# shared by Q/K/V/DO.
|
||||
stride_z, stride_h, stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M1: tl.constexpr, #
|
||||
BLOCK_N1: tl.constexpr, #
|
||||
BLOCK_M2: tl.constexpr, #
|
||||
BLOCK_N2: tl.constexpr, #
|
||||
HEAD_DIM: tl.constexpr):
|
||||
LN2 = 0.6931471824645996 # = ln(2)
|
||||
|
||||
bhid = tl.program_id(2)
|
||||
off_chz = (bhid * N_CTX).to(tl.int64)
|
||||
adj = (stride_h * (bhid % H) + stride_z * (bhid // H)).to(tl.int64)
|
||||
pid = tl.program_id(0)
|
||||
|
||||
# offset pointers for batch/head
|
||||
Q += adj
|
||||
K += adj
|
||||
V += adj
|
||||
DO += adj
|
||||
DQ += adj
|
||||
DK += adj
|
||||
DV += adj
|
||||
M += off_chz
|
||||
D += off_chz
|
||||
|
||||
# load scales
|
||||
offs_k = tl.arange(0, HEAD_DIM)
|
||||
|
||||
start_n = pid * BLOCK_N1
|
||||
start_m = 0
|
||||
|
||||
offs_n = start_n + tl.arange(0, BLOCK_N1)
|
||||
|
||||
dv = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
|
||||
dk = tl.zeros([BLOCK_N1, HEAD_DIM], dtype=tl.float32)
|
||||
|
||||
# load K and V: they stay in SRAM throughout the inner loop.
|
||||
k = tl.load(K + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
v = tl.load(V + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
|
||||
|
||||
num_steps = N_CTX // BLOCK_M1
|
||||
|
||||
dk, dv = _attn_bwd_dkdv( #
|
||||
dk, dv, #
|
||||
Q, k, v, sm_scale, #
|
||||
DO, #
|
||||
M, D, #
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M1, BLOCK_N1, HEAD_DIM, #
|
||||
start_n, start_m, num_steps #
|
||||
)
|
||||
|
||||
dv_ptrs = DV + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
tl.store(dv_ptrs, dv)
|
||||
|
||||
# Write back dK.
|
||||
dk *= sm_scale
|
||||
dk_ptrs = DK + offs_n[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
tl.store(dk_ptrs, dk)
|
||||
|
||||
# THIS BLOCK DOES DQ:
|
||||
start_m = pid * BLOCK_M2
|
||||
end_n = 0
|
||||
|
||||
offs_m = start_m + tl.arange(0, BLOCK_M2)
|
||||
|
||||
q = tl.load(Q + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
dq = tl.zeros([BLOCK_M2, HEAD_DIM], dtype=tl.float32)
|
||||
do = tl.load(DO + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d)
|
||||
|
||||
m = tl.load(M + offs_m)
|
||||
m = m[:, None]
|
||||
|
||||
num_steps = N_CTX // BLOCK_N2
|
||||
dq = _attn_bwd_dq(dq, q, K, V, #
|
||||
do, m, D, #
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
stride_tok, stride_d, #
|
||||
H, N_CTX, #
|
||||
BLOCK_M2, BLOCK_N2, HEAD_DIM, #
|
||||
start_m, end_n, num_steps #
|
||||
)
|
||||
# Write back dQ.
|
||||
dq_ptrs = DQ + offs_m[:, None] * stride_tok + offs_k[None, :] * stride_d
|
||||
dq *= LN2
|
||||
tl.store(dq_ptrs, dq)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
# ──────────────────────────── SPARSE ADDITION BEGIN ───────────────────────────
|
||||
def triton_block_sparse_attn_forward(q, k, v, q2k_index, q2k_num, variable_block_sizes):
|
||||
B, H, T, D = q.shape
|
||||
sm_scale = 1.0 / math.sqrt(D)
|
||||
max_kv_blks = q2k_index.shape[-1]
|
||||
assert T % 64 == 0, f"T must be a multiple of 64, but got {T}"
|
||||
assert T // 64 == q2k_num.shape[-1], f"shape mismatch, T // 64 = {T // 64}, q2k_num.shape[-2] = {q2k_num.shape[-2]}"
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((B, H, T), dtype=torch.float32, device=q.device)
|
||||
|
||||
grid = lambda _: (triton.cdiv(T, 64), B * H, 1)
|
||||
_attn_fwd_sparse[grid](
|
||||
q, k, v, sm_scale,
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
variable_block_sizes,
|
||||
M, o,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3),
|
||||
k.stride(0), k.stride(1), k.stride(2), k.stride(3),
|
||||
v.stride(0), v.stride(1), v.stride(2), v.stride(3),
|
||||
o.stride(0), o.stride(1), o.stride(2), o.stride(3),
|
||||
B, H, T,
|
||||
HEAD_DIM=D, STAGE=3
|
||||
)
|
||||
|
||||
return o, M
|
||||
|
||||
def triton_block_sparse_attn_backward(do, q, k, v, o, M, q2k_index, q2k_num, k2q_index, k2q_num, variable_block_sizes):
|
||||
assert do.is_contiguous()
|
||||
assert q.stride() == k.stride() == v.stride() == o.stride() == do.stride()
|
||||
|
||||
B, H, T, D = q.shape
|
||||
sm_scale = 1.0 / math.sqrt(D)
|
||||
dq = torch.empty_like(q)
|
||||
dk = torch.empty_like(k)
|
||||
dv = torch.empty_like(v)
|
||||
BATCH, N_HEAD, N_CTX = q.shape[:3]
|
||||
BLOCK_M1, BLOCK_N1, BLOCK_M2, BLOCK_N2 = 32, 64, 64, 32
|
||||
RCP_LN2 = 1.4426950408889634 # = 1.0 / ln(2)
|
||||
arg_k = k
|
||||
arg_k = arg_k * (sm_scale * RCP_LN2)
|
||||
PRE_BLOCK = 64
|
||||
assert N_CTX % PRE_BLOCK == 0
|
||||
pre_grid = (N_CTX // PRE_BLOCK, BATCH * N_HEAD)
|
||||
delta = torch.empty_like(M)
|
||||
_attn_bwd_preprocess[pre_grid](
|
||||
o, do, #
|
||||
delta, #
|
||||
BATCH, N_HEAD, N_CTX, #
|
||||
BLOCK_M=PRE_BLOCK, HEAD_DIM=D #
|
||||
)
|
||||
|
||||
|
||||
max_q_blks = k2q_index.shape[-1]
|
||||
max_kv_blks = q2k_index.shape[-1]
|
||||
|
||||
grid = (N_CTX // BLOCK_N1, 1, BATCH * N_HEAD)
|
||||
_attn_bwd[grid](
|
||||
q, arg_k, v, sm_scale, do, dq, dk, dv, #
|
||||
M, delta, #
|
||||
q2k_index, q2k_num, max_kv_blks,
|
||||
k2q_index, k2q_num, max_q_blks,
|
||||
variable_block_sizes,
|
||||
q.stride(0), q.stride(1), q.stride(2), q.stride(3), #
|
||||
N_HEAD, N_CTX, #
|
||||
BLOCK_M1=BLOCK_M1, BLOCK_N1=BLOCK_N1, #
|
||||
BLOCK_M2=BLOCK_M2, BLOCK_N2=BLOCK_N2, #
|
||||
HEAD_DIM=D #
|
||||
)
|
||||
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
@@ -3,6 +3,8 @@
|
||||
#include "kittens.cuh"
|
||||
#include <cooperative_groups.h>
|
||||
#include <iostream>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
|
||||
using namespace kittens;
|
||||
namespace cg = cooperative_groups;
|
||||
@@ -44,225 +46,9 @@ template<int D> struct fwd_globals {
|
||||
|
||||
int32_t *__restrict__ q2k_block_sparse_index;
|
||||
int32_t *__restrict__ q2k_block_sparse_num;
|
||||
int32_t *__restrict__ block_size;
|
||||
};
|
||||
|
||||
template<int D>
|
||||
__global__ __launch_bounds__(128, 3) // encourage compiler to reduce register usage so that an SM can hold 3 CTAs. Performance will drop from 391T to 353T if not specified explicitly.
|
||||
void fwd_attend_ker_even(const __grid_constant__ fwd_globals<D> g) {
|
||||
extern __shared__ int __shm[];
|
||||
tma_swizzle_allocator al((int*)&__shm[0]);
|
||||
|
||||
using K = fwd_attend_ker_tile_dims<D>;
|
||||
|
||||
using q_tile = st_bf<64, K::tile_width>;
|
||||
using k_tile = st_bf<128, K::tile_width>;
|
||||
using v_tile = st_bf<128, K::tile_width>;
|
||||
using k_tile_half = st_bf<64, K::tile_width>;
|
||||
using v_tile_half = st_bf<64, K::tile_width>;
|
||||
using l_col_vec = col_vec<st_fl<64, K::tile_width>>;
|
||||
using o_tile = st_bf<64, K::tile_width>;
|
||||
|
||||
q_tile (&q_smem)[1] = al.allocate<q_tile, 1>();
|
||||
|
||||
k_tile_half (&k_smem_0)[1] = al.allocate<k_tile_half, 1 >();
|
||||
k_tile_half (&k_smem_1)[1] = al.allocate<k_tile_half, 1 >();
|
||||
k_tile (*k_smem) = reinterpret_cast<k_tile(*)>(k_smem_0);
|
||||
|
||||
v_tile_half (&v_smem_0)[1] = al.allocate<v_tile_half, 1 >();
|
||||
v_tile_half (&v_smem_1)[1] = al.allocate<v_tile_half, 1 >();
|
||||
v_tile (*v_smem) = reinterpret_cast<v_tile(*)>(v_smem_0);
|
||||
|
||||
l_col_vec (&l_smem)[1] = al.allocate<l_col_vec, 1>();
|
||||
|
||||
auto (*o_smem) = reinterpret_cast<o_tile(*)>(q_smem);
|
||||
|
||||
int kv_head_idx = blockIdx.y / g.hr;
|
||||
int seq_idx = blockIdx.x;
|
||||
|
||||
|
||||
int32_t* q2k_block_sparse_index_ptr = g.q2k_block_sparse_index + blockIdx.z * gridDim.y * gridDim.x * g.max_kv_blocks_per_q + blockIdx.y * gridDim.x * g.max_kv_blocks_per_q + blockIdx.x * g.max_kv_blocks_per_q;
|
||||
int32_t* q2k_block_sparse_num_ptr = g.q2k_block_sparse_num + blockIdx.z * gridDim.y * gridDim.x + blockIdx.y * gridDim.x + blockIdx.x;
|
||||
int32_t kv_blocks = q2k_block_sparse_num_ptr[0] / 2; // each iter load 2 kv blocks
|
||||
__shared__ kittens::semaphore qsmem_semaphore, k_smem_arrived, v_smem_arrived;
|
||||
if (threadIdx.x == 0) {
|
||||
int32_t kv_block_index[2];
|
||||
reinterpret_cast<float2*>(kv_block_index)[0] = reinterpret_cast<float2*>(q2k_block_sparse_index_ptr)[0];
|
||||
|
||||
init_semaphore(qsmem_semaphore, 0, 1);
|
||||
init_semaphore(k_smem_arrived, 0, 1);
|
||||
init_semaphore(v_smem_arrived, 0, 1);
|
||||
|
||||
// preload q block
|
||||
coord<q_tile> q_tile_idx = {blockIdx.z, blockIdx.y, seq_idx, 0};
|
||||
tma::expect_bytes(qsmem_semaphore, sizeof(q_smem));
|
||||
tma::load_async(q_smem[0], g.q, q_tile_idx, qsmem_semaphore);
|
||||
|
||||
// preload the zeroth block of kv
|
||||
tma::expect_bytes(k_smem_arrived, sizeof(k_tile));
|
||||
coord<k_tile_half> k_tile_idx_0 = {blockIdx.z, kv_head_idx, kv_block_index[0], 0};
|
||||
coord<k_tile_half> k_tile_idx_1 = {blockIdx.z, kv_head_idx, kv_block_index[1], 0};
|
||||
tma::load_async(k_smem_0[0], g.k, k_tile_idx_0, k_smem_arrived);
|
||||
tma::load_async(k_smem_1[0], g.k, k_tile_idx_1, k_smem_arrived);
|
||||
|
||||
tma::expect_bytes(v_smem_arrived, sizeof(v_tile));
|
||||
coord<v_tile_half> v_tile_idx_0 = {blockIdx.z, kv_head_idx, kv_block_index[0], 0};
|
||||
coord<v_tile_half> v_tile_idx_1 = {blockIdx.z, kv_head_idx, kv_block_index[1], 0};
|
||||
tma::load_async(v_smem_0[0], g.v, v_tile_idx_0, v_smem_arrived);
|
||||
tma::load_async(v_smem_1[0], g.v, v_tile_idx_1, v_smem_arrived);
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
rt_fl<16, 128> att_block;
|
||||
rt_bf<16, 128> att_block_mma;
|
||||
rt_fl<16, K::tile_width> o_reg;
|
||||
|
||||
col_vec<rt_fl<16, 128>> max_vec, norm_vec, max_vec_last_scaled, max_vec_scaled;
|
||||
|
||||
neg_infty(max_vec);
|
||||
zero(norm_vec);
|
||||
zero(o_reg);
|
||||
|
||||
// wait for q block
|
||||
wait(qsmem_semaphore, 0);
|
||||
|
||||
for (int kv_idx = 0; kv_idx < kv_blocks - 1; kv_idx++) {
|
||||
// preload kv index
|
||||
int32_t kv_block_index[2];
|
||||
reinterpret_cast<float2*>(kv_block_index)[0] = reinterpret_cast<float2*>(q2k_block_sparse_index_ptr)[kv_idx + 1];
|
||||
|
||||
// wait k
|
||||
wait(k_smem_arrived, kv_idx % 2);
|
||||
|
||||
// compute QK^T
|
||||
warpgroup::mm_ABt(att_block, q_smem[0], k_smem[0]);
|
||||
|
||||
copy(max_vec_last_scaled, max_vec);
|
||||
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
|
||||
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
|
||||
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
// load K
|
||||
if (threadIdx.x == 0) {
|
||||
tma::expect_bytes(k_smem_arrived, sizeof(k_tile));
|
||||
coord<k_tile_half> k_tile_idx_0 = {blockIdx.z, kv_head_idx, kv_block_index[0], 0};
|
||||
coord<k_tile_half> k_tile_idx_1 = {blockIdx.z, kv_head_idx, kv_block_index[1], 0};
|
||||
tma::load_async(k_smem_0[0], g.k, k_tile_idx_0, k_smem_arrived);
|
||||
tma::load_async(k_smem_1[0], g.k, k_tile_idx_1, k_smem_arrived);
|
||||
}
|
||||
|
||||
// exp
|
||||
row_max(max_vec, att_block, max_vec);
|
||||
|
||||
if constexpr (D == 64) {
|
||||
mul(att_block, att_block, 1.44269504089f*0.125f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
|
||||
}
|
||||
else {
|
||||
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
|
||||
}
|
||||
|
||||
sub_row(att_block, att_block, max_vec_scaled);
|
||||
exp2(att_block, att_block);
|
||||
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
|
||||
exp2(max_vec_last_scaled, max_vec_last_scaled);
|
||||
mul(norm_vec, norm_vec, max_vec_last_scaled);
|
||||
row_sum(norm_vec, att_block, norm_vec);
|
||||
add(att_block, att_block, 0.f);
|
||||
copy(att_block_mma, att_block);
|
||||
mul_row(o_reg, o_reg, max_vec_last_scaled);
|
||||
|
||||
// wait v
|
||||
wait(v_smem_arrived, kv_idx % 2);
|
||||
|
||||
// compute SV
|
||||
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[0]);
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
// load V
|
||||
if (threadIdx.x == 0) {
|
||||
tma::expect_bytes(v_smem_arrived, sizeof(v_tile));
|
||||
// coord<v_tile> v_tile_idx = {blockIdx.z, kv_head_idx, kv_idx + 1, 0};
|
||||
// tma::load_async(v_smem[0], g.v, v_tile_idx, v_smem_arrived);
|
||||
coord<v_tile_half> v_tile_idx_0 = {blockIdx.z, kv_head_idx, kv_block_index[0], 0};
|
||||
coord<v_tile_half> v_tile_idx_1 = {blockIdx.z, kv_head_idx, kv_block_index[1], 0};
|
||||
tma::load_async(v_smem_0[0], g.v, v_tile_idx_0, v_smem_arrived);
|
||||
tma::load_async(v_smem_1[0], g.v, v_tile_idx_1, v_smem_arrived);
|
||||
}
|
||||
}
|
||||
|
||||
// last iter
|
||||
{
|
||||
int kv_idx = kv_blocks - 1;
|
||||
// wait k
|
||||
wait(k_smem_arrived, kv_idx % 2);
|
||||
|
||||
// compute QK^T
|
||||
warpgroup::mm_ABt(att_block, q_smem[0], k_smem[0]);
|
||||
|
||||
copy(max_vec_last_scaled, max_vec);
|
||||
if constexpr (D == 64) { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.125f); }
|
||||
else { mul(max_vec_last_scaled, max_vec_last_scaled, 1.44269504089f*0.08838834764f); }
|
||||
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
// exp
|
||||
row_max(max_vec, att_block, max_vec);
|
||||
|
||||
if constexpr (D == 64) {
|
||||
mul(att_block, att_block, 1.44269504089f*0.125f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.125f);
|
||||
}
|
||||
else {
|
||||
mul(att_block, att_block, 1.44269504089f*0.08838834764f);
|
||||
mul(max_vec_scaled, max_vec, 1.44269504089f*0.08838834764f);
|
||||
}
|
||||
|
||||
sub_row(att_block, att_block, max_vec_scaled);
|
||||
exp2(att_block, att_block);
|
||||
sub(max_vec_last_scaled, max_vec_last_scaled, max_vec_scaled);
|
||||
exp2(max_vec_last_scaled, max_vec_last_scaled);
|
||||
mul(norm_vec, norm_vec, max_vec_last_scaled);
|
||||
row_sum(norm_vec, att_block, norm_vec);
|
||||
add(att_block, att_block, 0.f);
|
||||
copy(att_block_mma, att_block);
|
||||
mul_row(o_reg, o_reg, max_vec_last_scaled);
|
||||
|
||||
// wait v
|
||||
wait(v_smem_arrived, kv_idx % 2);
|
||||
|
||||
// compute SV
|
||||
warpgroup::mma_AB(o_reg, att_block_mma, v_smem[0]);
|
||||
warpgroup::mma_async_wait();
|
||||
}
|
||||
|
||||
div_row(o_reg, o_reg, norm_vec);
|
||||
warpgroup::store(o_smem[0], o_reg);
|
||||
__syncthreads();
|
||||
|
||||
// TK store_async internally calls syncwarp so we need to route on warp level
|
||||
if (threadIdx.x / 32 == 0) {
|
||||
coord<o_tile> o_tile_idx = {blockIdx.z, blockIdx.y, seq_idx, 0};
|
||||
tma::store_async(g.o, o_smem[0], o_tile_idx);
|
||||
}
|
||||
|
||||
mul(max_vec_scaled, max_vec_scaled, 0.69314718056f);
|
||||
log(norm_vec, norm_vec);
|
||||
add(norm_vec, norm_vec, max_vec_scaled);
|
||||
|
||||
if constexpr (D == 64) { mul(norm_vec, norm_vec, -8.0f); }
|
||||
else { mul(norm_vec, norm_vec, -11.313708499f); }
|
||||
|
||||
warpgroup::store(l_smem[0], norm_vec);
|
||||
__syncthreads();
|
||||
|
||||
if (threadIdx.x / 32 == 0) {
|
||||
coord<l_col_vec> tile_idx = {blockIdx.z, blockIdx.y, 0, seq_idx};
|
||||
tma::store_async(g.l, l_smem[0], tile_idx);
|
||||
}
|
||||
tma::store_async_wait();
|
||||
}
|
||||
|
||||
template<int D>
|
||||
__global__ __launch_bounds__(128, 4)
|
||||
@@ -355,6 +141,7 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) { // use block siz
|
||||
}
|
||||
|
||||
// exp
|
||||
right_fill(att_block, att_block, g.block_size[q2k_block_sparse_index_ptr[kv_idx]], base_types::constants<float>::neg_infty());
|
||||
row_max(max_vec, att_block, max_vec);
|
||||
|
||||
if constexpr (D == 64) {
|
||||
@@ -407,6 +194,8 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) { // use block siz
|
||||
warpgroup::mma_async_wait();
|
||||
|
||||
// exp
|
||||
right_fill(att_block, att_block, g.block_size[q2k_block_sparse_index_ptr[kv_idx]], base_types::constants<float>::neg_infty());
|
||||
|
||||
row_max(max_vec, att_block, max_vec);
|
||||
|
||||
if constexpr (D == 64) {
|
||||
@@ -482,8 +271,9 @@ struct bwd_prep_globals {
|
||||
d_gl d;
|
||||
};
|
||||
|
||||
constexpr int PREP_NUM_WARPS = (1);
|
||||
template<int D>
|
||||
__global__ __launch_bounds__(4*kittens::WARP_THREADS, (D == 64) ? 2 : 1)
|
||||
__global__ __launch_bounds__(PREP_NUM_WARPS*kittens::WARP_THREADS, (D == 64) ? 6 / PREP_NUM_WARPS : 3 / PREP_NUM_WARPS)
|
||||
void bwd_attend_prep_ker(const __grid_constant__ bwd_prep_globals<D> g) {
|
||||
extern __shared__ int __shm[];
|
||||
tma_swizzle_allocator al((int*)&__shm[0]);
|
||||
@@ -494,9 +284,9 @@ void bwd_attend_prep_ker(const __grid_constant__ bwd_prep_globals<D> g) {
|
||||
using o_tile = st_bf<4*16, D>;
|
||||
using d_tile = col_vec<st_fl<4*16, D>>;
|
||||
|
||||
og_tile (&og_smem)[4] = al.allocate<og_tile, 4>();
|
||||
o_tile (&o_smem) [4] = al.allocate<o_tile , 4>();
|
||||
d_tile (&d_smem) [4] = al.allocate<d_tile , 4>();
|
||||
og_tile (&og_smem)[PREP_NUM_WARPS] = al.allocate<og_tile, PREP_NUM_WARPS>();
|
||||
o_tile (&o_smem) [PREP_NUM_WARPS] = al.allocate<o_tile , PREP_NUM_WARPS>();
|
||||
d_tile (&d_smem) [PREP_NUM_WARPS] = al.allocate<d_tile , PREP_NUM_WARPS>();
|
||||
|
||||
rt_fl<4*16, D> og_reg, o_reg;
|
||||
col_vec<rt_fl<4*16, D>> d_reg;
|
||||
@@ -505,13 +295,13 @@ void bwd_attend_prep_ker(const __grid_constant__ bwd_prep_globals<D> g) {
|
||||
|
||||
if (threadIdx.x == 0) {
|
||||
init_semaphore(smem_semaphore, 0, 1);
|
||||
tma::expect_bytes(smem_semaphore, sizeof(og_smem[0]) * 4 * 2);
|
||||
tma::expect_bytes(smem_semaphore, sizeof(og_smem[0]) * PREP_NUM_WARPS * 2);
|
||||
}
|
||||
__syncthreads();
|
||||
|
||||
if (warpid == 0) {
|
||||
for (int w = 0; w < 4; w++) {
|
||||
coord<o_tile> tile_idx = {blockIdx.z, blockIdx.y, (blockIdx.x * 4) + w, 0};
|
||||
for (int w = 0; w < PREP_NUM_WARPS; w++) {
|
||||
coord<o_tile> tile_idx = {blockIdx.z, blockIdx.y, (blockIdx.x * PREP_NUM_WARPS) + w, 0};
|
||||
tma::load_async(o_smem[w], g.o, tile_idx, smem_semaphore);
|
||||
tma::load_async(og_smem[w], g.og, tile_idx, smem_semaphore);
|
||||
}
|
||||
@@ -526,8 +316,8 @@ void bwd_attend_prep_ker(const __grid_constant__ bwd_prep_globals<D> g) {
|
||||
__syncthreads();
|
||||
|
||||
if (warpid == 0) {
|
||||
for (int w = 0; w < 4; w++) {
|
||||
coord<d_tile> tile_idx = {blockIdx.z, blockIdx.y, 0, (blockIdx.x * 4) + w};
|
||||
for (int w = 0; w < PREP_NUM_WARPS; w++) {
|
||||
coord<d_tile> tile_idx = {blockIdx.z, blockIdx.y, 0, (blockIdx.x * PREP_NUM_WARPS) + w};
|
||||
tma::store_async(g.d, d_smem[w], tile_idx);
|
||||
}
|
||||
}
|
||||
@@ -589,6 +379,7 @@ struct bwd_globals {
|
||||
|
||||
int32_t *__restrict__ k2q_block_sparse_index;
|
||||
int32_t *__restrict__ k2q_block_sparse_num;
|
||||
int32_t *__restrict__ block_size;
|
||||
};
|
||||
|
||||
__device__ static inline void
|
||||
@@ -711,7 +502,7 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
|
||||
|
||||
// wait for kv
|
||||
wait(kv_b, 0);
|
||||
|
||||
int fill_start = g.block_size[blockIdx.x] - 16 * kittens::warpid();
|
||||
for (int qo_idx = 0; qo_idx < qo_blocks - 1; qo_idx++) {
|
||||
// preload q index
|
||||
store_qg_block_index = load_q_block_index;
|
||||
@@ -730,7 +521,8 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
|
||||
|
||||
if constexpr (D == 64) { mul(s_block_t, s_block_t, 1.44269504089f*0.125f); }
|
||||
else { mul(s_block_t, s_block_t, 1.44269504089f*0.08838834764f); }
|
||||
|
||||
|
||||
lower_fill(s_block_t, s_block_t, fill_start, base_types::constants<float>::neg_infty());
|
||||
exp2(s_block_t, s_block_t); // P_i
|
||||
copy(p_block_t, s_block_t);
|
||||
copy(p_block_t_mma, s_block_t);
|
||||
@@ -776,7 +568,6 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
|
||||
__syncthreads(); // wait for sd_smem shared memory write
|
||||
warpgroup::mm_AtB(qg_reg, ds_smem_t[0], k_smem[0]); //delat dQ = dSK
|
||||
warpgroup::mma_commit_group();
|
||||
tma::store_async_wait();
|
||||
warpgroup::mma_async_wait();
|
||||
// store qg to shared memory
|
||||
warpgroup::store(qg_smem, qg_reg);
|
||||
@@ -786,6 +577,7 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
|
||||
if (threadIdx.x / 32 == 0) {
|
||||
coord<qg_tile> tile_idx = {blockIdx.z, blockIdx.y, store_qg_block_index, 0};
|
||||
tma::store_add_async(g.qg, qg_smem, tile_idx);
|
||||
tma::store_async_wait();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -808,7 +600,7 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
|
||||
|
||||
if constexpr (D == 64) { mul(s_block_t, s_block_t, 1.44269504089f*0.125f); }
|
||||
else { mul(s_block_t, s_block_t, 1.44269504089f*0.08838834764f); }
|
||||
|
||||
lower_fill(s_block_t, s_block_t, fill_start, base_types::constants<float>::neg_infty());
|
||||
exp2(s_block_t, s_block_t); // P_i
|
||||
copy(p_block_t, s_block_t);
|
||||
copy(p_block_t_mma, s_block_t);
|
||||
@@ -832,7 +624,6 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
|
||||
__syncthreads(); // wait for sd_smem shared memory write
|
||||
warpgroup::mm_AtB(qg_reg, ds_smem_t[0], k_smem[0]); //delat dQ = dSK
|
||||
warpgroup::mma_commit_group();
|
||||
tma::store_async_wait();
|
||||
warpgroup::mma_async_wait();
|
||||
// store qg to shared memory
|
||||
warpgroup::store(qg_smem, qg_reg);
|
||||
@@ -842,13 +633,14 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
|
||||
if (threadIdx.x / 32 == 0) {
|
||||
coord<qg_tile> tile_idx = {blockIdx.z, blockIdx.y, store_qg_block_index, 0};
|
||||
tma::store_add_async(g.qg, qg_smem, tile_idx);
|
||||
tma::store_async_wait();
|
||||
}
|
||||
}
|
||||
|
||||
// store kq and vq
|
||||
|
||||
// ! the following two line seems unnecessary.
|
||||
tma::store_async_wait(); // ensure qg is finished
|
||||
// tma::store_async_wait(); // ensure qg is finished
|
||||
__syncthreads();
|
||||
|
||||
warpgroup::store(kg_smem[0], kg_reg);
|
||||
@@ -874,7 +666,14 @@ void bwd_attend_ker(const __grid_constant__ bwd_globals<D> g) {
|
||||
#include <iostream>
|
||||
|
||||
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)
|
||||
block_sparse_attention_forward(
|
||||
torch::Tensor q,
|
||||
torch::Tensor k,
|
||||
torch::Tensor v,
|
||||
torch::Tensor q2k_block_sparse_index,
|
||||
torch::Tensor q2k_block_sparse_num,
|
||||
torch::Tensor block_size
|
||||
)
|
||||
{
|
||||
CHECK_INPUT(q);
|
||||
CHECK_INPUT(k);
|
||||
@@ -886,6 +685,10 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
|
||||
auto qo_heads = q.size(1);
|
||||
auto kv_heads = k.size(1);
|
||||
auto max_kv_blocks_per_q = q2k_block_sparse_index.size(3);
|
||||
auto num_q_blocks = block_size.size(0);
|
||||
TORCH_CHECK(batch==1, "Batch size dim will be removed in the future, please set batch to 1");
|
||||
TORCH_CHECK(num_q_blocks * 64 == seq_len, "This kernel supports variable block size, but it assumes the input sequence is properly padded.");
|
||||
TORCH_CHECK(num_q_blocks == q2k_block_sparse_index.size(2), "Number of Q blocks does not match between q2k_block_sparse_index and block_size");
|
||||
// check to see that these dimensions match for all inputs
|
||||
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
|
||||
@@ -940,8 +743,9 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
|
||||
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 == 64) {
|
||||
using q_tile = st_bf<fwd_attend_ker_tile_dims<64>::qo_height, fwd_attend_ker_tile_dims<64>::tile_width>;
|
||||
@@ -964,9 +768,21 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
|
||||
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
|
||||
|
||||
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(hr), static_cast<int>(max_kv_blocks_per_q), reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()), reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr())};
|
||||
globals g{
|
||||
qg_arg,
|
||||
kg_arg,
|
||||
vg_arg,
|
||||
lg_arg,
|
||||
og_arg,
|
||||
static_cast<int>(seq_len),
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_kv_blocks_per_q),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())
|
||||
};
|
||||
|
||||
auto mem_size = 54000;
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
dim3 grid(seq_len/(64), qo_heads, batch);
|
||||
|
||||
@@ -979,7 +795,7 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
|
||||
fwd_attend_ker<64><<<grid, (128), mem_size, stream>>>(g);
|
||||
|
||||
CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
cudaStreamSynchronize(stream);
|
||||
// cudaStreamSynchronize(stream);
|
||||
}
|
||||
|
||||
if (head_dim == 128) {
|
||||
@@ -1003,9 +819,21 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
|
||||
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
|
||||
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
|
||||
|
||||
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(hr), static_cast<int>(max_kv_blocks_per_q), reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()), reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr())};
|
||||
globals g{
|
||||
qg_arg,
|
||||
kg_arg,
|
||||
vg_arg,
|
||||
lg_arg,
|
||||
og_arg,
|
||||
static_cast<int>(seq_len),
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_kv_blocks_per_q),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())
|
||||
};
|
||||
|
||||
auto mem_size = 54000;
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
dim3 grid(seq_len/(64), qo_heads, batch);
|
||||
|
||||
@@ -1018,11 +846,11 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
|
||||
fwd_attend_ker<128><<<grid, (128), mem_size, stream>>>(g);
|
||||
|
||||
CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
cudaStreamSynchronize(stream);
|
||||
// cudaStreamSynchronize(stream);
|
||||
}
|
||||
|
||||
return {o, l_vec};
|
||||
cudaDeviceSynchronize();
|
||||
//cudadevicesynchronize();
|
||||
}
|
||||
|
||||
std::vector<torch::Tensor>
|
||||
@@ -1033,7 +861,8 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
torch::Tensor l_vec,
|
||||
torch::Tensor og,
|
||||
torch::Tensor k2q_block_sparse_index,
|
||||
torch::Tensor k2q_block_sparse_num)
|
||||
torch::Tensor k2q_block_sparse_num,
|
||||
torch::Tensor block_size)
|
||||
{
|
||||
CHECK_INPUT(q);
|
||||
CHECK_INPUT(k);
|
||||
@@ -1046,7 +875,7 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
auto seq_len = q.size(2);
|
||||
auto head_dim = q.size(3);
|
||||
auto max_q_blocks_per_kv = k2q_block_sparse_index.size(3);
|
||||
|
||||
TORCH_CHECK(k2q_block_sparse_index.size(2) == block_size.size(0), "k2q_block_sparse_index.size(2) must match block_size.size(0)");
|
||||
// check to see that these dimensions match for all inputs
|
||||
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
|
||||
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
|
||||
@@ -1132,16 +961,17 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
float* d_kg = reinterpret_cast<float*>(kg_ptr);
|
||||
float* d_vg = reinterpret_cast<float*>(vg_ptr);
|
||||
|
||||
auto mem_size = kittens::MAX_SHARED_MEMORY;
|
||||
auto threads = 4 * kittens::WARP_THREADS;
|
||||
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
|
||||
int threads = PREP_NUM_WARPS * kittens::WARP_THREADS;
|
||||
|
||||
cudaDeviceSynchronize();
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
//cudadevicesynchronize();
|
||||
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
cudaStreamSynchronize(stream);
|
||||
// cudaStreamSynchronize(stream);
|
||||
|
||||
// TORCH_CHECK(seq_len % (4*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 256");
|
||||
dim3 grid_bwd(seq_len/(4*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
dim3 grid_bwd(seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
|
||||
if (head_dim == 64) {
|
||||
using og_tile = st_bf<4*16, 64>;
|
||||
@@ -1216,13 +1046,13 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_q_blocks_per_kv),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr())
|
||||
};
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())};
|
||||
|
||||
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
|
||||
threads = 128;
|
||||
|
||||
cudaDeviceSynchronize();
|
||||
//cudadevicesynchronize();
|
||||
|
||||
{
|
||||
cudaFuncSetAttribute(
|
||||
@@ -1240,8 +1070,8 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
}
|
||||
|
||||
// CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
cudaStreamSynchronize(stream);
|
||||
cudaDeviceSynchronize();
|
||||
// cudaStreamSynchronize(stream);
|
||||
//cudadevicesynchronize();
|
||||
// const auto kernel_end = std::chrono::high_resolution_clock::now();
|
||||
// std::cout << "Kernel Time: " << std::chrono::duration_cast<std::chrono::microseconds>(kernel_end - start).count() << "us" << std::endl;
|
||||
// std::cout << "---" << std::endl;
|
||||
@@ -1320,13 +1150,13 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
static_cast<int>(hr),
|
||||
static_cast<int>(max_q_blocks_per_kv),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_index.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr())
|
||||
};
|
||||
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr()),
|
||||
reinterpret_cast<int32_t*>(block_size.data_ptr())};
|
||||
|
||||
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
|
||||
threads = 128;
|
||||
|
||||
cudaDeviceSynchronize();
|
||||
//cudadevicesynchronize();
|
||||
|
||||
{
|
||||
cudaFuncSetAttribute(
|
||||
@@ -1338,10 +1168,10 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
bwd_attend_ker<128><<<grid_bwd_2, threads, 113000, stream>>>(bwd_global);
|
||||
}
|
||||
|
||||
cudaStreamSynchronize(stream);
|
||||
cudaDeviceSynchronize();
|
||||
// cudaStreamSynchronize(stream);
|
||||
//cudadevicesynchronize();
|
||||
}
|
||||
|
||||
return {qg, kg, vg};
|
||||
cudaDeviceSynchronize();
|
||||
}
|
||||
//cudadevicesynchronize();
|
||||
}
|
||||
@@ -0,0 +1,185 @@
|
||||
import torch
|
||||
try:
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
except ImportError:
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
from vsa.block_sparse_attn_triton import triton_block_sparse_attn_forward, triton_block_sparse_attn_backward
|
||||
assert torch.__version__ >= "2.4.0", "VSA requires PyTorch 2.4.0 or higher"
|
||||
from vsa.index import map_to_index
|
||||
from typing import Tuple, Optional
|
||||
|
||||
|
||||
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_triton", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_triton(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
block_map = block_map.int()
|
||||
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
|
||||
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
|
||||
return o, M
|
||||
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_triton")
|
||||
def _block_sparse_attn_triton_fake(
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q = q.contiguous()
|
||||
k = k.contiguous()
|
||||
v = v.contiguous()
|
||||
o = torch.empty_like(q)
|
||||
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
|
||||
return o, M
|
||||
|
||||
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_backward_triton", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_backward_triton(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
M: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
|
||||
dq, dk, dv = triton_block_sparse_attn_backward(grad_output_padded, q_padded, k_padded, v_padded, o_padded, M, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes)
|
||||
return dq, dk, dv
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_backward_triton")
|
||||
def _block_sparse_attn_backward_triton_fake(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
M: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
dq = torch.empty_like(grad_output_padded)
|
||||
dk = torch.empty_like(grad_output_padded)
|
||||
dv = torch.empty_like(grad_output_padded)
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def backward_triton(ctx, grad_output1, grad_output2):
|
||||
q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes = ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_triton(grad_output1, q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
|
||||
def setup_context_triton(ctx, inputs, output):
|
||||
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
|
||||
o_padded, M = output
|
||||
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, M, block_map, variable_block_sizes)
|
||||
|
||||
block_sparse_attn_triton.register_autograd(backward_triton, setup_context=setup_context_triton)
|
||||
|
||||
|
||||
major, minor = torch.cuda.get_device_capability(0)
|
||||
|
||||
if major == 9 and minor == 0:# check if H100
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_SM90", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_SM90(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
)-> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q_padded = q_padded.contiguous()
|
||||
k_padded = k_padded.contiguous()
|
||||
v_padded = v_padded.contiguous()
|
||||
q2k_block_sparse_index, q2k_block_sparse_num = map_to_index(block_map)
|
||||
variable_block_sizes = variable_block_sizes.int()
|
||||
o_padded, lse_padded = block_sparse_fwd(q_padded, k_padded, v_padded, q2k_block_sparse_index, q2k_block_sparse_num, variable_block_sizes)
|
||||
return o_padded, lse_padded
|
||||
|
||||
|
||||
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_SM90")
|
||||
def _block_sparse_attn_SM90_fake(
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
q_padded, k_padded, v_padded = [x.contiguous() for x in (q_padded, k_padded, v_padded)]
|
||||
B, H, S, D = q_padded.shape
|
||||
o_padded = torch.empty_like(q_padded)
|
||||
lse_padded = torch.empty((B, H, S, 1), device=q_padded.device, dtype=torch.float32)
|
||||
return o_padded, lse_padded
|
||||
|
||||
|
||||
@torch.library.custom_op("vsa::block_sparse_attn_backward_SM90", mutates_args=(), device_types="cuda")
|
||||
def block_sparse_attn_backward_SM90(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
)-> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(block_map.transpose(-1, -2))
|
||||
grad_q_padded, grad_k_padded, grad_v_padded = block_sparse_bwd(
|
||||
q_padded, k_padded, v_padded, o_padded, lse_padded, grad_output_padded, k2q_block_sparse_index, k2q_block_sparse_num, variable_block_sizes
|
||||
)
|
||||
grad_q_padded = grad_q_padded.to(grad_output_padded.dtype)
|
||||
grad_k_padded = grad_k_padded.to(grad_output_padded.dtype)
|
||||
grad_v_padded = grad_v_padded.to(grad_output_padded.dtype)
|
||||
return grad_q_padded, grad_k_padded, grad_v_padded
|
||||
|
||||
@torch.library.register_fake("vsa::block_sparse_attn_backward_SM90")
|
||||
def _block_sparse_attn_backward_SM90_fake(
|
||||
grad_output_padded: torch.Tensor,
|
||||
q_padded: torch.Tensor,
|
||||
k_padded: torch.Tensor,
|
||||
v_padded: torch.Tensor,
|
||||
o_padded: torch.Tensor,
|
||||
lse_padded: torch.Tensor,
|
||||
block_map: torch.Tensor,
|
||||
variable_block_sizes: torch.Tensor,
|
||||
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
torch._check(grad_output_padded.dtype == torch.bfloat16)
|
||||
torch._check(lse_padded.dtype == torch.float32)
|
||||
grad_output_padded = grad_output_padded.contiguous()
|
||||
dq = torch.empty_like(grad_output_padded)
|
||||
dk = torch.empty_like(grad_output_padded)
|
||||
dv = torch.empty_like(grad_output_padded)
|
||||
return dq, dk, dv
|
||||
|
||||
|
||||
def backward_SM90(ctx, grad_output1, grad_output2):
|
||||
q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes= ctx.saved_tensors
|
||||
dq, dk, dv = block_sparse_attn_backward_SM90(grad_output1, q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
|
||||
return dq, dk, dv, None, None
|
||||
|
||||
def setup_context_SM90(ctx, inputs, output):
|
||||
q_padded, k_padded, v_padded, block_map, variable_block_sizes = inputs
|
||||
o_padded, lse_padded = output
|
||||
ctx.save_for_backward(q_padded, k_padded, v_padded, o_padded, lse_padded, block_map, variable_block_sizes)
|
||||
|
||||
|
||||
block_sparse_attn_SM90.register_autograd(backward_SM90, setup_context=setup_context_SM90)
|
||||
@@ -0,0 +1,152 @@
|
||||
|
||||
## pytorch sdpa version of block sparse ##
|
||||
import triton
|
||||
import triton.language as tl
|
||||
import torch
|
||||
|
||||
@triton.jit
|
||||
def topk_index_to_map_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
topk,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
for i in tl.static_range(topk):
|
||||
index = tl.load(index_ptr_base + i * index_kv_stride)
|
||||
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
|
||||
|
||||
@triton.jit
|
||||
def map_to_index_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
index_num_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
index_num_bs_stride,
|
||||
index_num_h_stride,
|
||||
index_num_q_stride,
|
||||
num_kv_blocks,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
num = 0
|
||||
for i in tl.range(num_kv_blocks):
|
||||
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
|
||||
if map_entry:
|
||||
tl.store(index_ptr_base + num * index_kv_stride, i)
|
||||
num += 1
|
||||
|
||||
tl.store(
|
||||
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
|
||||
q * index_num_q_stride, num)
|
||||
|
||||
def topk_index_to_map(index: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
transpose_map: bool = False):
|
||||
"""
|
||||
Convert topk indices to a map.
|
||||
|
||||
Args:
|
||||
index: [bs, h, num_q_blocks, topk]
|
||||
The topk indices tensor.
|
||||
num_kv_blocks: int
|
||||
The number of key-value blocks in the block_map returned
|
||||
transpose_map: bool
|
||||
If True, the block_map will be transposed on the final two dimensions.
|
||||
|
||||
Returns:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
A binary map where 1 indicates that the q block attends to the kv block.
|
||||
"""
|
||||
bs, h, num_q_blocks, topk = index.shape
|
||||
|
||||
if transpose_map is False:
|
||||
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
else:
|
||||
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
block_map = block_map.transpose(2, 3)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
topk_index_to_map_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
topk=topk,
|
||||
)
|
||||
|
||||
return block_map
|
||||
|
||||
def map_to_index(block_map: torch.Tensor):
|
||||
"""
|
||||
Convert a block map to indices and counts.
|
||||
|
||||
Args:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The block map tensor.
|
||||
|
||||
Returns:
|
||||
index: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The indices of the blocks.
|
||||
index_num: [bs, h, num_q_blocks]
|
||||
The number of blocks for each q block.
|
||||
"""
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
|
||||
|
||||
index = torch.full((block_map.shape),
|
||||
-1,
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
index_num = torch.empty((bs, h, num_q_blocks),
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
map_to_index_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
index_num,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
index_num.stride(0),
|
||||
index_num.stride(1),
|
||||
index_num.stride(2),
|
||||
num_kv_blocks=num_kv_blocks,
|
||||
)
|
||||
|
||||
return index, index_num
|
||||
@@ -1,470 +0,0 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
from torch.utils.checkpoint import detach_variable
|
||||
from typing import Tuple
|
||||
try:
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
except ImportError:
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
|
||||
|
||||
BLOCK_M = 64
|
||||
BLOCK_N = 64
|
||||
|
||||
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 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 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
|
||||
|
||||
@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]
|
||||
"""
|
||||
return BlockSparseAttentionFunction.apply(
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
)
|
||||
|
||||
def block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num):
|
||||
"""
|
||||
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks].
|
||||
[*, *, i, j] = 1 means the i-th q block should attend to the j-th kv block.
|
||||
"""
|
||||
# assert all elements in q2k_block_sparse_num can be devisible by 2
|
||||
o, lse = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
return o, lse
|
||||
|
||||
def block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
grad_output = grad_output.contiguous()
|
||||
grad_q, grad_k, grad_v = block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
return grad_q, grad_k, grad_v
|
||||
|
||||
## pytorch sdpa version of block sparse ##
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
@triton.jit
|
||||
def index_to_mask_kernel(
|
||||
q2k_block_sparse_index_ptr,
|
||||
q2k_block_sparse_num_ptr,
|
||||
mask_ptr,
|
||||
batch_size: tl.constexpr,
|
||||
num_heads: tl.constexpr,
|
||||
num_q_blocks: tl.constexpr,
|
||||
num_k_blocks: tl.constexpr,
|
||||
max_kv_blocks: tl.constexpr,
|
||||
BLOCK_Q: tl.constexpr,
|
||||
BLOCK_K: tl.constexpr,
|
||||
):
|
||||
bh, q, id = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64), tl.program_id(2).to(tl.int64)
|
||||
b = bh // num_heads
|
||||
h = bh % num_heads
|
||||
|
||||
num_valid_blocks = tl.load(q2k_block_sparse_num_ptr + b * num_heads * num_q_blocks + h * num_q_blocks + q)
|
||||
|
||||
if num_valid_blocks <= id:
|
||||
return
|
||||
k = tl.load(q2k_block_sparse_index_ptr + b * num_heads * num_q_blocks * max_kv_blocks + h * num_q_blocks * max_kv_blocks + q * max_kv_blocks + id)
|
||||
|
||||
full_mask = (tl.arange(0, BLOCK_Q)[:, None] < BLOCK_Q) & (tl.arange(0, BLOCK_K)[None, :] < BLOCK_K)
|
||||
|
||||
q_lengths = num_q_blocks * BLOCK_Q
|
||||
k_lengths = num_k_blocks * BLOCK_K
|
||||
mask_ptr_base = mask_ptr + b * num_heads * q_lengths * k_lengths + h * q_lengths * k_lengths + q * BLOCK_Q * k_lengths + k * BLOCK_K
|
||||
|
||||
tl.store(mask_ptr_base + tl.arange(0, BLOCK_Q)[:, None] * k_lengths + tl.arange(0, BLOCK_K)[None, :], full_mask)
|
||||
|
||||
def index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, BLOCK_Q, BLOCK_K, num_k_blocks):
|
||||
"""
|
||||
Convert block sparse indices to a mask.
|
||||
|
||||
Args:
|
||||
q2k_block_sparse_index: Indices for query-to-key sparse blocks
|
||||
q2k_block_sparse_num: Number of sparse blocks for each query block
|
||||
|
||||
Returns:
|
||||
mask: Block sparse mask tensor
|
||||
"""
|
||||
batch_size, num_heads, num_q_blocks, max_kv_blocks = q2k_block_sparse_index.shape
|
||||
assert q2k_block_sparse_num.shape == (batch_size, num_heads, num_q_blocks)
|
||||
|
||||
mask = torch.zeros((batch_size, num_heads, num_q_blocks * BLOCK_Q, num_k_blocks * BLOCK_K), dtype=torch.bool, device=q2k_block_sparse_index.device)
|
||||
|
||||
grid = (batch_size * num_heads, num_q_blocks, max_kv_blocks)
|
||||
index_to_mask_kernel[grid](
|
||||
q2k_block_sparse_index,
|
||||
q2k_block_sparse_num,
|
||||
mask,
|
||||
batch_size,
|
||||
num_heads,
|
||||
num_q_blocks,
|
||||
num_k_blocks,
|
||||
max_kv_blocks,
|
||||
BLOCK_Q=BLOCK_Q,
|
||||
BLOCK_K=BLOCK_K,
|
||||
)
|
||||
|
||||
return mask
|
||||
|
||||
@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
|
||||
|
||||
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_attention_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_attention_backward(
|
||||
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
|
||||
|
||||
|
||||
class DummyOperator(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, x):
|
||||
return x
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
return grad_output
|
||||
|
||||
class CheckpointSDPA(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, obj, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k):
|
||||
"""Forward pass."""
|
||||
with torch.no_grad():
|
||||
mask = index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k, k.shape[2] // block_k)
|
||||
outputs = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
|
||||
ctx.save_for_backward(*detach_variable((q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)))
|
||||
ctx.block_q = block_q
|
||||
ctx.block_k = block_k
|
||||
# the obj is passed in, then it can access the saved input
|
||||
# tensors later for recomputation
|
||||
obj.ctx = ctx
|
||||
return outputs
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
"""Backward pass."""
|
||||
inputs = ctx.saved_tensors
|
||||
output = ctx.output
|
||||
torch.autograd.backward(output, grad_output)
|
||||
ctx.output = None
|
||||
grads = tuple(inp.grad for inp in inputs)
|
||||
return (None, ) + grads + (None, None)
|
||||
|
||||
|
||||
class BlockSparseAttnTorch:
|
||||
def __init__(self):
|
||||
self.ctx = None
|
||||
|
||||
def recompute_mask(self, _):
|
||||
recomputed_mask = index_to_mask(self.q2k_block_sparse_index, self.q2k_block_sparse_num, self.block_q, self.block_k, self.num_kv_blocks)
|
||||
mask_size = recomputed_mask.untyped_storage().size()
|
||||
self.mask.untyped_storage().resize_(mask_size)
|
||||
self.mask.untyped_storage().copy_(recomputed_mask.untyped_storage())
|
||||
|
||||
def recompute(self, _):
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num = self.ctx.saved_tensors
|
||||
block_q = self.ctx.block_q
|
||||
block_k = self.ctx.block_k
|
||||
mask = index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k, k.shape[2] // block_k)
|
||||
with torch.enable_grad():
|
||||
output = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
|
||||
self.ctx.output = output
|
||||
self.ctx = None
|
||||
|
||||
@torch._dynamo.disable
|
||||
def forward(self, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k):
|
||||
"""
|
||||
Differentiable block sparse attention function using PyTorch.
|
||||
|
||||
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
|
||||
block_q: Block size for query
|
||||
block_k: Block size for key-value
|
||||
|
||||
Returns:
|
||||
output: Attention output tensor [batch_size, num_heads, seq_len_q, head_dim]
|
||||
"""
|
||||
|
||||
output = CheckpointSDPA.apply(
|
||||
self, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k
|
||||
)
|
||||
|
||||
o = DummyOperator.apply(output)
|
||||
o.register_hook(self.recompute)
|
||||
return o
|
||||
@@ -1,7 +1,9 @@
|
||||
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
|
||||
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
@@ -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.10.0 -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.10 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
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/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.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/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
EXPOSE 22
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
|
||||
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
@@ -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.11.11 -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.11 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
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/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.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/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
EXPOSE 22
|
||||
|
||||
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
|
||||
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
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
@@ -58,15 +58,15 @@ RUN source $HOME/.local/bin/env && \
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_sta.py install
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_vsa.py install
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
EXPOSE 22
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
FROM nvidia/cuda:12.9.1-cudnn-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
wget \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
zsh \
|
||||
vim \
|
||||
curl \
|
||||
gcc-11 \
|
||||
g++-11 \
|
||||
clang-11 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Set up C++20 compilers for ThunderKittens
|
||||
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
# Set CUDA environment variables
|
||||
ENV CUDA_HOME=/usr/local/cuda-12.9
|
||||
ENV PATH=${CUDA_HOME}/bin:${PATH}
|
||||
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
# Install uv and source its environment
|
||||
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
|
||||
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject.toml ./
|
||||
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
# Create and activate virtual environment with specific Python version and seed
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.12 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/sliding_tile_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn/video_sparse_attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup.py install
|
||||
|
||||
EXPOSE 22
|
||||
@@ -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
|
||||
|
After Width: | Height: | Size: 194 KiB |
@@ -9,11 +9,11 @@
|
||||
## Initialization Configuration
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.v1.configs.pipelines.PipelineConfig
|
||||
fastvideo.configs.pipelines.PipelineConfig
|
||||
```
|
||||
|
||||
## Sampling Configuration
|
||||
|
||||
```{autodoc2-summary}
|
||||
fastvideo.v1.configs.sample.SamplingParam
|
||||
fastvideo.configs.sample.SamplingParam
|
||||
```
|
||||
|
||||
@@ -18,7 +18,6 @@ import os
|
||||
import re
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import requests
|
||||
|
||||
@@ -97,8 +96,7 @@ copybutton_prompt_is_regexp = True
|
||||
#
|
||||
html_title = project
|
||||
html_theme = 'sphinx_book_theme'
|
||||
html_logo = '../../assets/logo.jpg'
|
||||
#html_favicon = 'assets/logos/vllm-logo-only-light.ico'
|
||||
html_logo = '../../assets/logos/icon_simple.svg'
|
||||
html_theme_options = {
|
||||
'path_to_docs': 'docs/source',
|
||||
'repository_url': 'https://github.com/hao-ai-lab/FastVideo/',
|
||||
@@ -168,8 +166,7 @@ _cached_base: str = ""
|
||||
_cached_branch: str = ""
|
||||
|
||||
|
||||
def get_repo_base_and_branch(
|
||||
pr_number: str) -> tuple[Optional[str], Optional[str]]:
|
||||
def get_repo_base_and_branch(pr_number: str) -> tuple[str | None, str | None]:
|
||||
global _cached_base, _cached_branch
|
||||
if _cached_base and _cached_branch:
|
||||
return _cached_base, _cached_branch
|
||||
|
||||
@@ -3,12 +3,12 @@
|
||||
|
||||
If you prefer a containerized development environment or want to avoid managing dependencies manually, you can use our prebuilt Docker image:
|
||||
|
||||
**Image:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
|
||||
**Images:** [`ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest`](https://ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev)
|
||||
|
||||
## Starting the container
|
||||
|
||||
```bash
|
||||
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
|
||||
docker run --gpus all -it ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest
|
||||
```
|
||||
|
||||
This will:
|
||||
|
||||
@@ -6,14 +6,16 @@ You can easily use the FastVideo Docker image as a custom container on [RunPod](
|
||||
|
||||
## Creating a new pod
|
||||
|
||||
Choose a GPU that supports CUDA 12.4
|
||||
Choose a GPU that supports CUDA 12.8
|
||||
|
||||
Pick 1 or 2 L40S GPU(s)
|
||||
|
||||

|
||||
|
||||
When creating your pod template, use this image:
|
||||
|
||||
```
|
||||
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev: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)):
|
||||
|
||||
@@ -22,10 +22,20 @@ source ~/.bashrc
|
||||
Create and activate a Conda environment for FastVideo:
|
||||
|
||||
```
|
||||
conda create -n fastvideo python=3.10 -y
|
||||
conda create -n fastvideo python=3.12 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
Install `uv` (optional, but recommended):
|
||||
|
||||
From instructions on [uv](https://astral.sh/uv/):
|
||||
|
||||
```
|
||||
curl -LsSf https://astral.sh/uv/install.sh | sh
|
||||
# or
|
||||
wget -qO- https://astral.sh/uv/install.sh | sh
|
||||
```
|
||||
|
||||
Clone the FastVideo repository and go to the FastVideo directory:
|
||||
|
||||
```
|
||||
@@ -36,10 +46,10 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
Now you can install FastVideo and setup git hooks for running linting. By using `pre-commit`, the linters will run and have to pass before you'll be able to make a commit.
|
||||
|
||||
```bash
|
||||
pip install -e .[dev]
|
||||
uv pip install -e .[dev]
|
||||
|
||||
# Can also install flash-attn (optional)
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
uv pip install flash-attn --no-build-isolation
|
||||
|
||||
# Linting, formatting and static type checking
|
||||
pre-commit install --hook-type pre-commit --hook-type commit-msg
|
||||
@@ -50,3 +60,14 @@ pre-commit run --all-files
|
||||
# Unit tests
|
||||
pytest tests/
|
||||
```
|
||||
|
||||
If you are on a Hopper GPU, you should also install [FA3](https://github.com/Dao-AILab/flash-attention) for much better performance:
|
||||
|
||||
```
|
||||
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention/hopper
|
||||
|
||||
# make sure you have ninja installed
|
||||
uv pip install ninja
|
||||
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
@@ -1,24 +1,24 @@
|
||||
# 🔍 FastVideo Overview
|
||||
|
||||
This document outlines FastVideo's architecture for developers interested in framework internals or contributions. It serves as an onboarding guide for new contributors by providing an overview of the most important directories and files within the `fastvideo/v1/` codebase.
|
||||
This document outlines FastVideo's architecture for developers interested in framework internals or contributions. It serves as an onboarding guide for new contributors by providing an overview of the most important directories and files within the `fastvideo/` codebase.
|
||||
|
||||
## Table of Contents - V1 Directory Structure and Files
|
||||
## Table of Contents - Directory Structure and Files
|
||||
|
||||
- [`fastvideo/v1/pipelines/`](#design-pipeline-system) - Core diffusion pipeline components
|
||||
- [`fastvideo/v1/models/`](#design-model-components) - Model implementations
|
||||
- [`fastvideo/pipelines/`](#design-pipeline-system) - Core diffusion pipeline components
|
||||
- [`fastvideo/models/`](#design-model-components) - Model implementations
|
||||
- [`dits/`](#design-transformer-models) - Transformer-based diffusion models
|
||||
- [`vaes/`](#design-vae-variational-auto-encoder) - Variational autoencoders
|
||||
- [`encoders/`](#design-text-and-image-encoders) - Text and image encoders
|
||||
- [`schedulers/`](#design-schedulers) - Diffusion schedulers
|
||||
- [`fastvideo/v1/attention/`](#design-optimized-attention) - Optimized attention implementations
|
||||
- [`fastvideo/v1/distributed/`](#design-distributed-processing) - Distributed computing utilities
|
||||
- [`fastvideo/v1/layers/`](#design-tensor-parallelism) - Custom neural network layers
|
||||
- [`fastvideo/v1/platforms/`](#design-platforms) - Hardware platform abstractions
|
||||
- [`fastvideo/v1/worker/`](#design-executor-and-worker-abstractions) - Multi-GPU process management
|
||||
- [`fastvideo/v1/fastvideo_args.py`](#design-fastvideo-args) - Argument handling
|
||||
- [`fastvideo/v1/forward_context.py`](#design-forwardcontext) - Forward pass context management
|
||||
- `fastvideo/v1/utils.py` - Utility functions
|
||||
- [`fastvideo/v1/logger.py`](#design-logger) - Logging infrastructure
|
||||
- [`fastvideo/attention/`](#design-optimized-attention) - Optimized attention implementations
|
||||
- [`fastvideo/distributed/`](#design-distributed-processing) - Distributed computing utilities
|
||||
- [`fastvideo/layers/`](#design-tensor-parallelism) - Custom neural network layers
|
||||
- [`fastvideo/platforms/`](#design-platforms) - Hardware platform abstractions
|
||||
- [`fastvideo/worker/`](#design-executor-and-worker-abstractions) - Multi-GPU process management
|
||||
- [`fastvideo/fastvideo_args.py`](#design-fastvideo-args) - Argument handling
|
||||
- [`fastvideo/forward_context.py`](#design-forwardcontext) - Forward pass context management
|
||||
- `fastvideo/utils.py` - Utility functions
|
||||
- [`fastvideo/logger.py`](#design-logger) - Logging infrastructure
|
||||
|
||||
## Core Architecture
|
||||
|
||||
@@ -32,7 +32,7 @@ FastVideo separates model components from execution logic with these principles:
|
||||
(design-fastvideo-args)=
|
||||
## FastVideoArgs
|
||||
|
||||
The `FastVideoArgs` class in `fastvideo/v1/fastvideo_args.py` serves as the central configuration system for FastVideo. It contains all parameters needed to control model loading, inference configuration, performance optimization settings, and more.
|
||||
The `FastVideoArgs` class in `fastvideo/fastvideo_args.py` serves as the central configuration system for FastVideo. It contains all parameters needed to control model loading, inference configuration, performance optimization settings, and more.
|
||||
|
||||
Key features include:
|
||||
- **Command-line Interface**: Automatic conversion between CLI arguments and dataclass fields
|
||||
@@ -111,7 +111,7 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward
|
||||
(design-forwardbatch)=
|
||||
### ForwardBatch
|
||||
|
||||
Defined in `fastvideo/v1/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsulates the data payload passed between pipeline stages. It typically holds:
|
||||
Defined in `fastvideo/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsulates the data payload passed between pipeline stages. It typically holds:
|
||||
|
||||
- **Input Data**: Prompts, images, generation parameters
|
||||
- **Intermediate State**: Embeddings, latents, timesteps, accumulated during stage execution
|
||||
@@ -123,14 +123,14 @@ This structure facilitates clear state transitions between stages.
|
||||
(design-model-components)=
|
||||
## Model Components
|
||||
|
||||
The `fastvideo/v1/models/` directory contains implementations of the core neural network models used in video diffusion:
|
||||
The `fastvideo/models/` directory contains implementations of the core neural network models used in video diffusion:
|
||||
|
||||
(design-transformer-models)=
|
||||
### Transformer Models
|
||||
|
||||
Transformer networks perform the actual denoising during diffusion:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/dits/`
|
||||
- **Location**: `fastvideo/models/dits/`
|
||||
- **Examples**:
|
||||
- `WanTransformer3DModel`
|
||||
- `HunyuanVideoTransformer3DModel`
|
||||
@@ -157,7 +157,7 @@ def forward(
|
||||
|
||||
VAEs handle conversion between pixel space and latent space:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/vaes/`
|
||||
- **Location**: `fastvideo/models/vaes/`
|
||||
- **Examples**:
|
||||
- `AutoencoderKLWan`
|
||||
- `AutoencoderKLHunyuanVideo`
|
||||
@@ -175,7 +175,7 @@ FastVideo's VAE implementations include:
|
||||
|
||||
Encoders process conditioning inputs into embeddings:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/encoders/`
|
||||
- **Location**: `fastvideo/models/encoders/`
|
||||
- **Text Encoders**:
|
||||
- `CLIPTextModel`
|
||||
- `LlamaModel`
|
||||
@@ -193,7 +193,7 @@ FastVideo implements optimizations such as:
|
||||
|
||||
Schedulers manage the diffusion sampling process:
|
||||
|
||||
- **Location**: `fastvideo/v1/models/schedulers/`
|
||||
- **Location**: `fastvideo/models/schedulers/`
|
||||
- **Examples**:
|
||||
- `UniPCMultistepScheduler`
|
||||
- `FlowMatchEulerDiscreteScheduler`
|
||||
@@ -219,7 +219,7 @@ def step(
|
||||
(design-optimized-attention)=
|
||||
## Optimized Attention
|
||||
|
||||
The `fastvideo/v1/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
|
||||
The `fastvideo/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
|
||||
|
||||
### Attention Backends
|
||||
Multiple implementations with automatic selection:
|
||||
@@ -248,7 +248,7 @@ Supports various patterns with memory optimization techniques:
|
||||
(design-distributed-processing)=
|
||||
## Distributed Processing
|
||||
|
||||
The `fastvideo/v1/distributed/` directory contains implementations for distributed model execution:
|
||||
The `fastvideo/distributed/` directory contains implementations for distributed model execution:
|
||||
|
||||
(design-tensor-parallelism)=
|
||||
### Tensor Parallelism
|
||||
@@ -260,7 +260,7 @@ Tensor parallelism splits model weights across devices:
|
||||
|
||||
```python
|
||||
# Tensor-parallel layers in a transformer block
|
||||
from fastvideo.v1.layers.linear import ColumnParallelLinear, RowParallelLinear
|
||||
from fastvideo.layers.linear import ColumnParallelLinear, RowParallelLinear
|
||||
|
||||
# Split along output dimension
|
||||
self.qkv_proj = ColumnParallelLinear(
|
||||
@@ -288,7 +288,7 @@ Sequence parallelism splits sequences across devices:
|
||||
|
||||
```python
|
||||
# Distributed attention for long sequences
|
||||
from fastvideo.v1.attention import DistributedAttention
|
||||
from fastvideo.attention import DistributedAttention
|
||||
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
@@ -312,7 +312,7 @@ Efficient communication primitives minimize distributed overhead:
|
||||
|
||||
### ForwardContext
|
||||
|
||||
Defined in `fastvideo/v1/forward_context.py`, `ForwardContext` manages execution-specific state *within* a forward pass, particularly for low-level optimizations. It is accessed via `get_forward_context()`.
|
||||
Defined in `fastvideo/forward_context.py`, `ForwardContext` manages execution-specific state *within* a forward pass, particularly for low-level optimizations. It is accessed via `get_forward_context()`.
|
||||
|
||||
- **Attention Metadata**: Configuration for optimized attention kernels (`attn_metadata`)
|
||||
- **Profiling Data**: Potential hooks for performance metrics collection
|
||||
@@ -333,7 +333,7 @@ with set_forward_context(current_timestep, attn_metadata, fastvideo_args):
|
||||
(design-executor-and-worker-abstractions)=
|
||||
## Executor and Worker System
|
||||
|
||||
The `fastvideo/v1/worker/` directory contains the distributed execution framework:
|
||||
The `fastvideo/worker/` directory contains the distributed execution framework:
|
||||
|
||||
### Executor Abstraction
|
||||
|
||||
@@ -360,7 +360,7 @@ This design allows FastVideo to efficiently utilize multiple GPUs while providin
|
||||
(design-platforms)=
|
||||
## Platforms
|
||||
|
||||
The `fastvideo/v1/platforms/` directory provides hardware platform abstractions that enable FastVideo to run efficiently on different hardware configurations:
|
||||
The `fastvideo/platforms/` directory provides hardware platform abstractions that enable FastVideo to run efficiently on different hardware configurations:
|
||||
|
||||
### Platform Abstraction
|
||||
|
||||
@@ -377,7 +377,7 @@ The primary components include:
|
||||
Usage example:
|
||||
|
||||
```python
|
||||
from fastvideo.v1.platforms import current_platform, _Backend
|
||||
from fastvideo.platforms import current_platform, _Backend
|
||||
|
||||
# Check hardware capabilities
|
||||
if current_platform.supports_backend(_Backend.FLASH_ATTN):
|
||||
@@ -398,9 +398,9 @@ See [PR](https://github.com/hao-ai-lab/FastVideo/pull/356)
|
||||
|
||||
If you're a new contributor, here are some common areas to explore:
|
||||
|
||||
1. **Adding a new model**: Implement new model types in the appropriate subdirectory of `fastvideo/v1/models/`
|
||||
1. **Adding a new model**: Implement new model types in the appropriate subdirectory of `fastvideo/models/`
|
||||
2. **Optimizing performance**: Look at attention implementations or memory management
|
||||
3. **Adding a new pipeline**: Create a new pipeline subclass in `fastvideo/v1/pipelines/`
|
||||
3. **Adding a new pipeline**: Create a new pipeline subclass in `fastvideo/pipelines/`
|
||||
4. **Hardware support**: Extend the `platforms` module for new hardware targets
|
||||
|
||||
When adding code, follow these practices:
|
||||
|
||||
@@ -0,0 +1,43 @@
|
||||
(v0-data-preprocess)=
|
||||
|
||||
# 🧱 Data Preprocess for Distillation
|
||||
|
||||
For distillation, we use the same data preprocessing pipeline as training. Please refer to the [Training Data Preprocess](../training/data_preprocess.md) for general preprocessing steps.
|
||||
|
||||
## Distillation-Specific Datasets
|
||||
|
||||
### FastVideo 480P Synthetic Wan Dataset
|
||||
|
||||
For Wan2.1 T2V distillation, we use the **FastVideo 480P Synthetic Wan dataset** ([FastVideo/Wan-Syn_77x448x832_600k](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k)) which contains 600k synthetic latents.
|
||||
|
||||
```bash
|
||||
# Download the preprocessed dataset
|
||||
python scripts/huggingface/download_hf.py \
|
||||
--repo_id "FastVideo/Wan-Syn_77x448x832_600k" \
|
||||
--local_dir "FastVideo/Wan-Syn_77x448x832_600k" \
|
||||
--repo_type "dataset"
|
||||
```
|
||||
|
||||
### Crush Smol Dataset
|
||||
|
||||
For Wan2.2 TI2V distillation, we use the crush_smol dataset which includes both raw videos and preprocessed latents.
|
||||
|
||||
```bash
|
||||
# Download dataset
|
||||
python scripts/huggingface/download_hf.py \
|
||||
--repo_id=FastVideo/mini_i2v_dataset \
|
||||
--local_dir=data/mini_i2v_dataset \
|
||||
--repo_type=dataset
|
||||
```
|
||||
|
||||
## Preprocessing for Distillation
|
||||
|
||||
The preprocessing steps are identical to training. Run the appropriate preprocessing script based on your model:
|
||||
|
||||
```bash
|
||||
# For Wan2.1 T2V
|
||||
bash scripts/preprocess/v1_preprocess_wan_data_t2v
|
||||
|
||||
# For Wan2.2 TI2V
|
||||
bash examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/preprocess_wan_data_ti2v_5b.sh
|
||||
```
|
||||
@@ -0,0 +1,87 @@
|
||||
# 🎯 Distillation
|
||||
|
||||
We introduce a new finetuning strategy - **Sparse-distill**, which jointly integrates **[DMD](https://arxiv.org/abs/2405.14867)** and **[VSA](https://arxiv.org/abs/2505.13389)** in a single training process. This approach combines the benefits of both distillation to shorten diffusion steps and sparse attention to reduce attention computations, enabling much faster video generation.
|
||||
|
||||
## 📊 Model Overview
|
||||
|
||||
We provide two distilled models:
|
||||
|
||||
- **[FastWan2.1-T2V-1.3B-Diffusers](https://huggingface.co/FastVideo/FastWan2.1-T2V-1.3B-Diffusers)**: 3-step inference, up to **16 FPS** on H100 GPU
|
||||
- **[FastWan2.1-T2V-14B-480P-Diffusers](https://huggingface.co/FastVideo/FastWan2.1-T2V-14B-480P-Diffusers)**: 3-step inference, up to **60x speed up** at 480P, **90x speed up** at 720P for denoising loop
|
||||
- **[FastWan2.2-TI2V-5B-FullAttn-Diffusers](https://huggingface.co/FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers)**: 3-step inference, up to **50x speed up** at 720P for denoising loop
|
||||
|
||||
Both models are trained on **61×448×832** resolution but support generating videos with **any resolution** (1.3B model mainly support 480P, 14B model support 480P and 720P, quality may degrade for different resolutions).
|
||||
|
||||
## ⚙️ Inference
|
||||
First install [VSA](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation.html). Set `MODEL_BASE` to your own model path and run:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_dmd.sh
|
||||
```
|
||||
|
||||
## 🗂️ Dataset
|
||||
|
||||
We use the **FastVideo 480P Synthetic Wan dataset** ([FastVideo/Wan-Syn_77x448x832_600k](https://huggingface.co/datasets/FastVideo/Wan-Syn_77x448x832_600k)) for distillation, which contains 600k synthetic latents.
|
||||
|
||||
### Download Dataset
|
||||
|
||||
```bash
|
||||
# Download the preprocessed dataset
|
||||
python scripts/huggingface/download_hf.py \
|
||||
--repo_id "FastVideo/Wan-Syn_77x448x832_600k" \
|
||||
--local_dir "FastVideo/Wan-Syn_77x448x832_600k" \
|
||||
--repo_type "dataset"
|
||||
```
|
||||
|
||||
## 🚀 Training Scripts
|
||||
|
||||
### Wan2.1 1.3B Model Sparse-Distill
|
||||
|
||||
For the 1.3B model, we use **4 nodes with 32 H200 GPUs** (8 GPUs per node):
|
||||
|
||||
```bash
|
||||
# Multi-node training (8 nodes, 64 GPUs total)
|
||||
sbatch examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P/distill_dmd_VSA_t2v_1.3B.slurm
|
||||
```
|
||||
|
||||
**Key Configuration:**
|
||||
- Global batch size: 64
|
||||
- Gradient accumulation steps: 2
|
||||
- Learning rate: 1e-5
|
||||
- VSA attention sparsity: 0.8
|
||||
- Training steps: 4000 (~12 hours)
|
||||
|
||||
### Wan2.1 14B Model Sparse-Distill
|
||||
|
||||
For the 14B model, we use **8 nodes with 64 H200 GPUs** (8 GPUs per node):
|
||||
|
||||
```bash
|
||||
# Multi-node training (8 nodes, 64 GPUs total)
|
||||
sbatch examples/distill/Wan2.1-T2V/Wan-Syn-Data-480P/distill_dmd_VSA_t2v_14B.slurm
|
||||
```
|
||||
|
||||
**Key Configuration:**
|
||||
- Global batch size: 64
|
||||
- Sequence parallel size: 4
|
||||
- Gradient accumulation steps: 4
|
||||
- Learning rate: 1e-5
|
||||
- VSA attention sparsity: 0.9
|
||||
- Training steps: 3000 (~52 hours)
|
||||
- HSDP shard dim: 8
|
||||
|
||||
### Wan2.2 5B Model Sparse-Distill
|
||||
|
||||
For the 5B model, we use **8 nodes with 64 H200 GPUs** (8 GPUs per node):
|
||||
|
||||
```bash
|
||||
# Multi-node training (8 nodes, 64 GPUs total)
|
||||
sbatch examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/distill_dmd_t2v_5B.sh
|
||||
```
|
||||
|
||||
**Key Configuration:**
|
||||
- Global batch size: 64
|
||||
- Sequence parallel size: 1
|
||||
- Gradient accumulation steps: 1
|
||||
- Learning rate: 2e-5
|
||||
- Training steps: 3000 (~12 hours)
|
||||
- HSDP shard dim: 1
|
||||
@@ -5,7 +5,6 @@ import itertools
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
ROOT_DIR = Path(__file__).parent.parent.parent.resolve()
|
||||
ROOT_DIR_RELATIVE = '../../../..'
|
||||
@@ -28,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
|
||||
}
|
||||
@@ -89,7 +97,7 @@ class Example:
|
||||
generate() -> str: Generates the documentation content.
|
||||
""" # noqa: E501
|
||||
path: Path
|
||||
category: Optional[str] = None
|
||||
category: str | None = None
|
||||
main_file: Path = field(init=False)
|
||||
other_files: list[Path] = field(init=False)
|
||||
title: str = field(init=False)
|
||||
@@ -162,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(
|
||||
@@ -194,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
|
||||
@@ -230,36 +268,197 @@ 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
|
||||
|
||||
if example.category == "training":
|
||||
# For training 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)
|
||||
|
||||
elif example.category == "distillation" and len(path_parts) >= 2:
|
||||
# For distillation examples like Wan2.1-T2V/Wan-Syn-Data-480P
|
||||
model = path_parts[0] # e.g., "Wan2.1-T2V"
|
||||
dataset = path_parts[1] # e.g., "Wan-Syn-Data-480P"
|
||||
method = "DMD" # Default method for distillation
|
||||
|
||||
# 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,119 @@
|
||||
# NVIDIA GPU
|
||||
|
||||
Instructions to install FastVideo for NVIDIA CUDA GPUs.
|
||||
|
||||
## Requirements
|
||||
|
||||
- **OS: Linux or Windows WSL**
|
||||
- **Python: 3.10-3.12**
|
||||
- **CUDA 12.8**
|
||||
- **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.
|
||||
Note that you can also use `uv` to install FastVideo in a Conda environment.
|
||||
:::
|
||||
|
||||
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 --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 --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.8 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,102 @@
|
||||
# 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.
|
||||
Note that you can also use `uv` to install FastVideo in a Conda environment.
|
||||
:::
|
||||
|
||||
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.
|
||||
@@ -10,20 +10,20 @@ This class will be the primary Python API for generating videos and images.
|
||||
fastvideo.VideoGenerator
|
||||
```
|
||||
|
||||
`````{py:class} VideoGenerator(fastvideo_args: fastvideo.v1.fastvideo_args.FastVideoArgs, executor_class: type[fastvideo.v1.worker.executor.Executor], log_stats: bool)
|
||||
:canonical: fastvideo.v1.entrypoints.video_generator.VideoGenerator
|
||||
`````{py:class} VideoGenerator(fastvideo_args: fastvideo.fastvideo_args.FastVideoArgs, executor_class: type[fastvideo.worker.executor.Executor], log_stats: bool)
|
||||
:canonical: fastvideo.entrypoints.video_generator.VideoGenerator
|
||||
|
||||
```{autodoc2-docstring} fastvideo.v1.entrypoints.video_generator.VideoGenerator
|
||||
```{autodoc2-docstring} fastvideo.entrypoints.video_generator.VideoGenerator
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
`VideoGenerator.from_pretrained()` should be the primary way of creating a new video generator.
|
||||
|
||||
````{py:method} from_pretrained(model_path: str, device: typing.Optional[str] = None, torch_dtype: typing.Optional[torch.dtype] = None, pipeline_config: typing.Optional[typing.Union[str | fastvideo.v1.configs.pipelines.PipelineConfig]] = None, **kwargs) -> fastvideo.v1.entrypoints.video_generator.VideoGenerator
|
||||
:canonical: fastvideo.v1.entrypoints.video_generator.VideoGenerator.from_pretrained
|
||||
````{py:method} from_pretrained(model_path: str, device: typing.Optional[str] = None, torch_dtype: typing.Optional[torch.dtype] = None, pipeline_config: typing.Optional[typing.Union[str | fastvideo.configs.pipelines.PipelineConfig]] = None, **kwargs) -> fastvideo.entrypoints.video_generator.VideoGenerator
|
||||
:canonical: fastvideo.entrypoints.video_generator.VideoGenerator.from_pretrained
|
||||
:classmethod:
|
||||
|
||||
```{autodoc2-docstring} fastvideo.v1.entrypoints.video_generator.VideoGenerator.from_pretrained
|
||||
```{autodoc2-docstring} fastvideo.entrypoints.video_generator.VideoGenerator.from_pretrained
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
@@ -38,25 +38,25 @@ The follow two classes `PipelineConfig` and `SamplingParam` are used to configur
|
||||
```
|
||||
|
||||
`````{py:class} PipelineConfig
|
||||
:canonical: fastvideo.v1.configs.pipelines.base.PipelineConfig
|
||||
:canonical: fastvideo.configs.pipelines.base.PipelineConfig
|
||||
|
||||
```{autodoc2-docstring} fastvideo.v1.configs.pipelines.base.PipelineConfig
|
||||
```{autodoc2-docstring} fastvideo.configs.pipelines.base.PipelineConfig
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
````{py:method} from_pretrained(model_path: str) -> fastvideo.v1.configs.pipelines.base.PipelineConfig
|
||||
:canonical: fastvideo.v1.configs.pipelines.base.PipelineConfig.from_pretrained
|
||||
````{py:method} from_pretrained(model_path: str) -> fastvideo.configs.pipelines.base.PipelineConfig
|
||||
:canonical: fastvideo.configs.pipelines.base.PipelineConfig.from_pretrained
|
||||
:classmethod:
|
||||
|
||||
```{autodoc2-docstring} fastvideo.v1.configs.pipelines.base.PipelineConfig.from_pretrained
|
||||
```{autodoc2-docstring} fastvideo.configs.pipelines.base.PipelineConfig.from_pretrained
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
|
||||
````{py:method} dump_to_json(file_path: str)
|
||||
:canonical: fastvideo.v1.configs.pipelines.base.PipelineConfig.dump_to_json
|
||||
:canonical: fastvideo.configs.pipelines.base.PipelineConfig.dump_to_json
|
||||
|
||||
```{autodoc2-docstring} fastvideo.v1.configs.pipelines.base.PipelineConfig.dump_to_json
|
||||
```{autodoc2-docstring} fastvideo.configs.pipelines.base.PipelineConfig.dump_to_json
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
@@ -68,16 +68,16 @@ The follow two classes `PipelineConfig` and `SamplingParam` are used to configur
|
||||
```
|
||||
|
||||
`````{py:class} SamplingParam
|
||||
:canonical: fastvideo.v1.configs.sample.base.SamplingParam
|
||||
:canonical: fastvideo.configs.sample.base.SamplingParam
|
||||
|
||||
```{autodoc2-docstring} fastvideo.v1.configs.sample.base.SamplingParam
|
||||
```{autodoc2-docstring} fastvideo.configs.sample.base.SamplingParam
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
````{py:method} from_pretrained(model_path: str) -> fastvideo.v1.configs.sample.base.SamplingParam
|
||||
:canonical: fastvideo.v1.configs.sample.base.SamplingParam.from_pretrained
|
||||
````{py:method} from_pretrained(model_path: str) -> fastvideo.configs.sample.base.SamplingParam
|
||||
:canonical: fastvideo.configs.sample.base.SamplingParam.from_pretrained
|
||||
:classmethod:
|
||||
|
||||
```{autodoc2-docstring} fastvideo.v1.configs.sample.base.SamplingParam.from_pretrained
|
||||
```{autodoc2-docstring} fastvideo.configs.sample.base.SamplingParam.from_pretrained
|
||||
:parser: docs.source.autodoc2_docstring_parser
|
||||
```
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# Welcome to FastVideo
|
||||
|
||||
:::{figure} ../../assets/logo.jpg
|
||||
:::{figure} ../../assets/logos/logo.svg
|
||||
:align: center
|
||||
:alt: FastVideo
|
||||
:class: no-scaled-link
|
||||
@@ -9,7 +9,7 @@
|
||||
|
||||
:::{raw} html
|
||||
<p style="text-align:center">
|
||||
<strong>FastVideo is a unified framework for accelerated video generation.
|
||||
<strong>FastVideo is a unified inference and post-training framework for accelerated video generation.
|
||||
</strong>
|
||||
</p>
|
||||
|
||||
@@ -21,11 +21,10 @@
|
||||
</p>
|
||||
:::
|
||||
|
||||
It features a clean, consistent API that works across popular video models, making it easier for developers to author new models and incorporate system- or kernel-level optimizations.
|
||||
With FastVideo's optimizations, you can achieve more than 3x inference improvement compared to other systems.
|
||||
FastVideo is an inference and post-training framework for diffusion models. It features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
|
||||
|
||||
<div style="text-align: center;">
|
||||
<img src=_static/images/perf.png width="100%"/>
|
||||
<img src=_static/images/fastwan.png width="100%"/>
|
||||
</div>
|
||||
|
||||
## Key Features
|
||||
@@ -35,16 +34,11 @@ FastVideo has the following features:
|
||||
- [Sliding Tile Attention](https://arxiv.org/pdf/2502.04507)
|
||||
- [TeaCache](https://arxiv.org/pdf/2411.19108)
|
||||
- [Sage Attention](https://arxiv.org/abs/2410.02367)
|
||||
- Cutting edge models
|
||||
- Wan2.1 T2V, I2V
|
||||
- HunyuanVideo
|
||||
- FastHunyuan: consistency distilled video diffusion models for 8x inference speedup.
|
||||
- StepVideo T2V
|
||||
- Distillation support
|
||||
- Recipes for video DiT, based on [PCM](https://github.com/G-U-N/Phased-Consistency-Model).
|
||||
- Support distilling/finetuning/inferencing state-of-the-art open video DiTs: 1. Mochi 2. Hunyuan.
|
||||
- Scalable training with FSDP, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
- Memory efficient finetuning with LoRA, precomputed latent, and precomputed text embeddings.
|
||||
- E2E post-training support
|
||||
- Data preprocessing pipeline for video data.
|
||||
- [Sparse distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/) for Wan2.1 and Wan2.2 using [Video Sparse Attention](https://arxiv.org/pdf/2505.13389) and [Distribution Matching Distillation](https://tianweiy.github.io/dmd2/)
|
||||
- Support full finetuning and LoRA finetuning for state-of-the-art open video DiTs.
|
||||
- Scalable training with FSDP2, sequence parallelism, and selective activation checkpointing, with near linear scaling to 64 GPUs.
|
||||
|
||||
## Documentation
|
||||
|
||||
@@ -63,22 +57,31 @@ 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
|
||||
:::
|
||||
|
||||
% What is STA Kernel?
|
||||
@@ -91,6 +94,15 @@ sliding_tile_attention/installation
|
||||
sliding_tile_attention/demo
|
||||
:::
|
||||
|
||||
% What is VSA Kernel?
|
||||
|
||||
:::{toctree}
|
||||
:caption: Video Sparse Attention
|
||||
:maxdepth: 1
|
||||
|
||||
video_sparse_attention/installation
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Design
|
||||
:maxdepth: 1
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -46,25 +46,25 @@ FastVideo uses the Hugging Face Diffusers format for model organization:
|
||||
### Implementing Modules
|
||||
|
||||
Place new modules in the appropriate directories:
|
||||
- Encoders: `fastvideo/v1/models/encoders/`
|
||||
- VAEs: `fastvideo/v1/models/vaes/`
|
||||
- Transformer models: `fastvideo/v1/models/dits/`
|
||||
- Schedulers: `fastvideo/v1/models/schedulers/`
|
||||
- Encoders: `fastvideo/models/encoders/`
|
||||
- VAEs: `fastvideo/models/vaes/`
|
||||
- Transformer models: `fastvideo/models/dits/`
|
||||
- Schedulers: `fastvideo/models/schedulers/`
|
||||
|
||||
### Adapting Model Layers
|
||||
|
||||
#### Layer Replacements
|
||||
Replace standard PyTorch layers with FastVideo optimized versions:
|
||||
- nn.LayerNorm → fastvideo.v1.layers.layernorm.RMSNorm
|
||||
- Embedding layers → fastvideo.v1.layers.vocab_parallel_embedding modules
|
||||
- Activation functions → versions from fastvideo.v1.layers.activation
|
||||
- nn.LayerNorm → fastvideo.layers.layernorm.RMSNorm
|
||||
- Embedding layers → fastvideo.layers.vocab_parallel_embedding modules
|
||||
- Activation functions → versions from fastvideo.layers.activation
|
||||
|
||||
#### Distributed Linear Layers
|
||||
Use appropriate parallel layers for distribution:
|
||||
|
||||
```python
|
||||
# Output dimension parallelism
|
||||
from fastvideo.v1.layers.linear import ColumnParallelLinear
|
||||
from fastvideo.layers.linear import ColumnParallelLinear
|
||||
self.q_proj = ColumnParallelLinear(
|
||||
input_size=hidden_size,
|
||||
output_size=head_size * num_heads,
|
||||
@@ -73,7 +73,7 @@ self.q_proj = ColumnParallelLinear(
|
||||
)
|
||||
|
||||
# Fused QKV projection
|
||||
from fastvideo.v1.layers.linear import QKVParallelLinear
|
||||
from fastvideo.layers.linear import QKVParallelLinear
|
||||
self.qkv_proj = QKVParallelLinear(
|
||||
hidden_size=hidden_size,
|
||||
head_size=attention_head_dim,
|
||||
@@ -82,7 +82,7 @@ self.qkv_proj = QKVParallelLinear(
|
||||
)
|
||||
|
||||
# Input dimension parallelism
|
||||
from fastvideo.v1.layers.linear import RowParallelLinear
|
||||
from fastvideo.layers.linear import RowParallelLinear
|
||||
self.out_proj = RowParallelLinear(
|
||||
input_size=head_size * num_heads,
|
||||
output_size=hidden_size,
|
||||
@@ -96,8 +96,8 @@ Replace standard attention with FastVideo's optimized attention:
|
||||
|
||||
```python
|
||||
# Local attention patterns
|
||||
from fastvideo.v1.attention import LocalAttention
|
||||
from fastvideo.v1.attention.backends.abstract import _Backend
|
||||
from fastvideo.attention import LocalAttention
|
||||
from fastvideo.attention.backends.abstract import _Backend
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
@@ -108,7 +108,7 @@ self.attn = LocalAttention(
|
||||
)
|
||||
|
||||
# Distributed attention for long sequences
|
||||
from fastvideo.v1.attention import DistributedAttention
|
||||
from fastvideo.attention import DistributedAttention
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
@@ -130,7 +130,7 @@ self.attn = DistributedAttention(
|
||||
Register implemented modules in the model registry:
|
||||
|
||||
```python
|
||||
# In fastvideo/v1/models/registry.py
|
||||
# In fastvideo/models/registry.py
|
||||
_TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"YourTransformerModel": ("dits", "yourmodule", "YourTransformerClass"),
|
||||
}
|
||||
@@ -145,7 +145,7 @@ _VAE_MODELS = {
|
||||
Create a new directory for your pipeline:
|
||||
|
||||
```
|
||||
fastvideo/v1/pipelines/
|
||||
fastvideo/pipelines/
|
||||
├── your_pipeline/
|
||||
│ ├── __init__.py
|
||||
│ └── your_pipeline.py
|
||||
@@ -167,13 +167,13 @@ Pipelines are composed of stages, each handling a specific part of the diffusion
|
||||
### Creating Your Pipeline
|
||||
|
||||
```python
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.stages import (
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.stages import (
|
||||
InputValidationStage, CLIPTextEncodingStage, TimestepPreparationStage,
|
||||
LatentPreparationStage, DenoisingStage, DecodingStage
|
||||
)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
import torch
|
||||
|
||||
class MyCustomPipeline(ComposedPipelineBase):
|
||||
@@ -246,7 +246,7 @@ EntryClass = MyCustomPipeline
|
||||
If existing stages don't meet your needs, create custom ones:
|
||||
|
||||
```python
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
|
||||
class MyCustomStage(PipelineStage):
|
||||
"""Custom processing stage for the pipeline."""
|
||||
@@ -305,7 +305,7 @@ EntryClass = [MyCustomPipeline, MyOtherPipeline]
|
||||
```
|
||||
|
||||
The registry will automatically:
|
||||
1. Scan all packages under `fastvideo/v1/pipelines/`
|
||||
1. Scan all packages under `fastvideo/pipelines/`
|
||||
2. Look for `EntryClass` variables
|
||||
3. Register pipelines using their class names as identifiers
|
||||
|
||||
|
||||
@@ -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.
|
||||
@@ -27,7 +27,7 @@ def main():
|
||||
model_name = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
config = PipelineConfig.from_pretrained(model_name)
|
||||
config.vae_precision = "fp16"
|
||||
config.use_cpu_offload = True
|
||||
config.dit_cpu_offload = True
|
||||
|
||||
# Create the generator
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
|
||||
@@ -5,7 +5,7 @@ This page contains step-by-step instructions to get you quickly started with vid
|
||||
## Requirements
|
||||
- **OS**: Linux (Tested on Ubuntu 22.04+)
|
||||
- **Python**: 3.10-3.12
|
||||
- **CUDA**: 12.4
|
||||
- **CUDA**: 12.8
|
||||
- **GPU**: At least one NVIDIA GPU
|
||||
|
||||
## Installation
|
||||
@@ -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).
|
||||
|
||||
@@ -19,6 +19,7 @@ This page describes the various options for speeding up generation times in Fast
|
||||
- Torch SDPA: `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`
|
||||
- Flash Attention 2 and 3: `FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN`
|
||||
- Sliding Tile Attention: `FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN`
|
||||
- Video Sparse Attention: `FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN`
|
||||
- Sage Attention: `FASTVIDEO_ATTENTION_BACKEND=SAGE_ATTN`
|
||||
|
||||
### Configuring Backends
|
||||
@@ -74,6 +75,17 @@ pip install st_attn==0.0.4
|
||||
|
||||
Please see [this page](#sta-installation) for more installation instructions.
|
||||
|
||||
(optimizations-vsa)=
|
||||
### Video Sparse Attention
|
||||
**`VIDEO_SPARSE_ATTN`**
|
||||
|
||||
```bash
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
Please see [this page](#vsa-installation) for more installation instructions.
|
||||
|
||||
(optimizations-sage)=
|
||||
### Sage Attention
|
||||
**`SAGE_ATTN`**
|
||||
|
||||
@@ -6,6 +6,7 @@ The symbols used have the following meanings:
|
||||
|
||||
- ✅ = Full compatibility
|
||||
- ❌ = No compatibility
|
||||
- ⭕ = Does not apply to this model
|
||||
|
||||
## Models x Optimization
|
||||
The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods and FastVideo will use the optimal default parameters when initializing and generating videos.
|
||||
@@ -37,51 +38,94 @@ The `HuggingFace Model ID` can be directly pass to `from_pretrained()` methods a
|
||||
* TeaCache
|
||||
* Sliding Tile Attn
|
||||
* Sage Attn
|
||||
* Video Sparse Attention (VSA)
|
||||
- * FastWan2.1 T2V 1.3B
|
||||
* `FastVideo/FastWan2.1-T2V-1.3B-Diffusers`
|
||||
* 480P
|
||||
* ⭕
|
||||
* ⭕
|
||||
* ⭕
|
||||
* ✅
|
||||
- * FastWan2.2 TI2V 5B Full Attn*
|
||||
* `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers`
|
||||
* 720P
|
||||
* ⭕
|
||||
* ⭕
|
||||
* ⭕
|
||||
* ✅
|
||||
- * Wan2.2 TI2V 5B
|
||||
* `Wan-AI/Wan2.2-TI2V-5B-Diffusers`
|
||||
* 720P
|
||||
* ⭕
|
||||
* ⭕
|
||||
* ✅
|
||||
* ⭕
|
||||
- * Wan2.2 T2V A14B
|
||||
* `Wan-AI/Wan2.2-T2V-A14B-Diffusers`
|
||||
* 480P<br>720P
|
||||
* ❌
|
||||
* ❌
|
||||
* ✅
|
||||
* ⭕
|
||||
- * Wan2.2 I2V A14B
|
||||
* `Wan-AI/Wan2.2-I2V-A14B-Diffusers`
|
||||
* 480P<br>720P
|
||||
* ❌
|
||||
* ❌
|
||||
* ✅
|
||||
* ⭕
|
||||
- * HunyuanVideo
|
||||
* `hunyuanvideo-community/HunyuanVideo`
|
||||
* 720px1280p<br>544px960p
|
||||
* ❌
|
||||
* ✅
|
||||
* ✅
|
||||
* ⭕
|
||||
- * FastHunyuan
|
||||
* `FastVideo/FastHunyuan-diffusers`
|
||||
* 720px1280p<br>544px960p
|
||||
* ❌
|
||||
* ✅
|
||||
* ✅
|
||||
- * Wan T2V 1.3B
|
||||
* ⭕
|
||||
- * Wan2.1 T2V 1.3B
|
||||
* `Wan-AI/Wan2.1-T2V-1.3B-Diffusers`
|
||||
* 480P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
- * Wan T2V 14B
|
||||
* ⭕
|
||||
- * Wan2.1 T2V 14B
|
||||
* `Wan-AI/Wan2.1-T2V-14B-Diffusers`
|
||||
* 480P, 720P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
- * Wan I2V 480P
|
||||
* ⭕
|
||||
- * Wan2.1 I2V 480P
|
||||
* `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers`
|
||||
* 480P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
- * Wan I2V 720P
|
||||
* ⭕
|
||||
- * Wan2.1 I2V 720P
|
||||
* `Wan-AI/Wan2.1-I2V-14B-720P-Diffusers`
|
||||
* 720P
|
||||
* ✅
|
||||
* ✅*
|
||||
* ✅
|
||||
* ✅
|
||||
* ⭕
|
||||
- * StepVideo T2V
|
||||
* `FastVideo/stepvideo-t2v-diffusers`
|
||||
* 768px768px204f<br>544px992px204f<br>544px992px136f
|
||||
* ❌
|
||||
* ❌
|
||||
* ✅
|
||||
* ⭕
|
||||
:::
|
||||
|
||||
**Note**: there are some known quality issues with Wan2.1 + Sliding Tile Attn. We are working on fixing this issue.
|
||||
**Note**: Wan2.2 TI2V 5B has some quality issues when performing I2V generation. We are working on fixing this issue.
|
||||
|
||||
## Special requirements
|
||||
|
||||
|
||||