Compare commits
83
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a6d4d815f8 | ||
|
|
42da276006 | ||
|
|
7b3c4fbdd1 | ||
|
|
f041bd3f49 | ||
|
|
fbb9a73ab4 | ||
|
|
daea023383 | ||
|
|
0422d18377 | ||
|
|
f46c6d923f | ||
|
|
79a31b0b45 | ||
|
|
ef98769b46 | ||
|
|
869c5f7370 | ||
|
|
74529f22a4 | ||
|
|
2a2d67e792 | ||
|
|
bb3367c9ac | ||
|
|
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 |
+65
-36
@@ -1,89 +1,122 @@
|
||||
env:
|
||||
IMAGE_VERSION: "py3.12-latest"
|
||||
BUILDKITE_CLEAN_CHECKOUT: true
|
||||
|
||||
steps:
|
||||
- label: "pre-commit"
|
||||
command: ".buildkite/scripts/pre_commit.sh"
|
||||
agents:
|
||||
queue: "default"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
|
||||
- wait
|
||||
|
||||
- label: "Trigger Tests"
|
||||
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/loader/**"
|
||||
- "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/loader/**"
|
||||
- "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/loader/**"
|
||||
- "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"
|
||||
- "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/v1/**"
|
||||
- "fastvideo/tests/lora/**"
|
||||
- "fastvideo/models/loader/**"
|
||||
- "fastvideo/tests/transformers/**"
|
||||
- "fastvideo/pipelines/**"
|
||||
- "fastvideo/layers/lora/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
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:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=training
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**"
|
||||
- "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/vsa/**"
|
||||
- "csrc/attn/tk/**"
|
||||
- "csrc/attn/setup_vsa.py"
|
||||
@@ -92,15 +125,14 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Training Tests VSA"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=training_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**"
|
||||
- "fastvideo/**"
|
||||
- "csrc/attn/st_attn/**"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/attn/config_sta.py"
|
||||
@@ -108,10 +140,9 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Inference Tests STA"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=inference_sta
|
||||
agents:
|
||||
queue: "default"
|
||||
@@ -123,10 +154,9 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests STA"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=precision_sta
|
||||
agents:
|
||||
queue: "default"
|
||||
@@ -139,10 +169,9 @@ steps:
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests VSA"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=precision_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
|
||||
@@ -50,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"
|
||||
@@ -58,7 +58,7 @@ if [ -z "${TEST_TYPE:-}" ]; then
|
||||
fi
|
||||
log "Test type: $TEST_TYPE"
|
||||
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT IMAGE_VERSION=$IMAGE_VERSION"
|
||||
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")
|
||||
@@ -81,6 +81,10 @@ case "$TEST_TYPE" in
|
||||
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"
|
||||
@@ -97,6 +101,14 @@ case "$TEST_TYPE" in
|
||||
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"
|
||||
exit 1
|
||||
|
||||
@@ -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
|
||||
@@ -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 @@
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
@@ -115,37 +115,37 @@ jobs:
|
||||
- 'csrc/attn/config_vsa.py'
|
||||
- 'csrc/attn/vsa.cpp'
|
||||
vsa-paths: &vsa-paths
|
||||
- 'fastvideo/v1/**'
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
|
||||
# Actual tests
|
||||
encoder-test:
|
||||
- 'fastvideo/v1/models/encoders/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/encoders/**'
|
||||
- 'fastvideo/models/encoders/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/encoders/**'
|
||||
- *common-paths
|
||||
vae-test:
|
||||
- 'fastvideo/v1/models/vaes/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/vaes/**'
|
||||
- 'fastvideo/models/vaes/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/vaes/**'
|
||||
- *common-paths
|
||||
transformer-test:
|
||||
- 'fastvideo/v1/models/dits/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/transformers/**'
|
||||
- 'fastvideo/v1/layers/**'
|
||||
- 'fastvideo/v1/attention/**'
|
||||
- 'fastvideo/models/dits/**'
|
||||
- 'fastvideo/models/loader/**'
|
||||
- 'fastvideo/tests/transformers/**'
|
||||
- 'fastvideo/layers/**'
|
||||
- 'fastvideo/attention/**'
|
||||
- *common-paths
|
||||
training-test:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
training-test-VSA:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
inference-test-STA:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'fastvideo/**'
|
||||
- *common-paths
|
||||
- *sta-kernel-paths
|
||||
precision-test-STA:
|
||||
@@ -167,7 +167,7 @@ jobs:
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/encoders -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -185,7 +185,7 @@ jobs:
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/vaes -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -203,7 +203,7 @@ jobs:
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/transformers -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -229,7 +229,7 @@ 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 }}
|
||||
@@ -248,7 +248,7 @@ jobs:
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/Vanilla -srP"
|
||||
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 }}
|
||||
@@ -264,11 +264,11 @@ jobs:
|
||||
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 }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/VSA -srP"
|
||||
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 }}
|
||||
@@ -284,11 +284,11 @@ jobs:
|
||||
with:
|
||||
job_id: "inference-test-STA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
gpu_count: 2
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/inference/STA -srP"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/tests/inference/STA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -326,7 +326,7 @@ jobs:
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_block_sparse.py"
|
||||
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 }}
|
||||
@@ -343,7 +343,7 @@ jobs:
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
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 }}
|
||||
@@ -372,4 +372,4 @@ jobs:
|
||||
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12", "training-test", "training-test-VSA", "inference-test-STA", "precision-test-STA", "precision-test-VSA"]'
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
@@ -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
-1
@@ -20,6 +20,7 @@ samples/
|
||||
data/
|
||||
outputs/
|
||||
outputs_video
|
||||
checkpoints/
|
||||
sbatch.sh
|
||||
*.out
|
||||
env
|
||||
@@ -40,6 +41,7 @@ eggs/
|
||||
docs/_build/
|
||||
docs/source/getting_started/examples/
|
||||
docs/source/inference/examples/
|
||||
docs/source/training/examples/
|
||||
|
||||
# VSCode
|
||||
.vscode/
|
||||
@@ -55,7 +57,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
|
||||
|
||||
@@ -3,7 +3,7 @@ default_stages:
|
||||
- manual # Run in CI
|
||||
exclude: |
|
||||
(?x)(
|
||||
fastvideo/v1/third_party/.*|
|
||||
fastvideo/third_party/.*|
|
||||
csrc/.*|
|
||||
assets/.*|
|
||||
tests/.*|
|
||||
@@ -60,7 +60,7 @@ repos:
|
||||
rev: v1.15.0
|
||||
hooks:
|
||||
- id: mypy
|
||||
args: [--python-version, '3.10', --follow-imports, "skip", ]
|
||||
args: [--python-version, '3.10', --follow-imports, "skip", "--disable-error-code", "union-attr", "--disable-error-code", "override" ]
|
||||
additional_dependencies: [types-cachetools, types-setuptools, types-PyYAML, types-requests]
|
||||
- repo: local
|
||||
hooks:
|
||||
@@ -69,7 +69,7 @@ repos:
|
||||
entry: bash
|
||||
args:
|
||||
- -c
|
||||
- 'git ls-files | grep -v "^fastvideo/v1/tests/ssim/" | grep " " && echo "Filenames should not contain spaces!" && exit 1 || exit 0'
|
||||
- 'git ls-files | grep -v "^fastvideo/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
|
||||
|
||||
@@ -8,13 +8,18 @@ It features a clean, consistent API that works across popular video models, maki
|
||||
With FastVideo's optimizations, you can achieve more than 3x inference improvement compared to other systems.
|
||||
|
||||
<p align="center">
|
||||
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg" target="_blank"> <b>Slack</b> </a> |
|
||||
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastHunyuan" target="_blank"><b>FastHunyuan</b></a> | 🤗 <a href="https://huggingface.co/FastVideo/FastMochi-diffusers" target="_blank"><b>FastMochi</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ" target="_blank"> <b>Slack</b> </a> |
|
||||
</p>
|
||||
|
||||
<div align="center">
|
||||
<img src=assets/perf.png width="90%"/>
|
||||
</div>
|
||||
|
||||
## NEWS
|
||||
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
- ```2025/02/18```: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
|
||||
## Key Features
|
||||
|
||||
FastVideo has the following features:
|
||||
@@ -125,9 +130,26 @@ We learned and reused code from the following projects:
|
||||
We thank MBZUAI and [Anyscale](https://www.anyscale.com/) for their support throughout this project.
|
||||
|
||||
## Citation
|
||||
If you use FastVideo for your research, please cite our paper:
|
||||
If you use FastVideo for your research, please cite our work:
|
||||
|
||||
```bibtex
|
||||
@software{fastvideo2024,
|
||||
title = {FastVideo: A Unified Framework for Accelerated Video Generation},
|
||||
author = {The FastVideo Team},
|
||||
url = {https://github.com/hao-ai-lab/FastVideo},
|
||||
month = apr,
|
||||
year = {2024},
|
||||
}
|
||||
|
||||
@misc{zhang2025vsafastervideodiffusion,
|
||||
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
|
||||
author={Peiyuan Zhang and Haofeng Huang and Yongqi Chen and Will Lin and Zhengzhong Liu and Ion Stoica and Eric Xing and Hao Zhang},
|
||||
year={2025},
|
||||
eprint={2505.13389},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV},
|
||||
url={https://arxiv.org/abs/2505.13389},
|
||||
}
|
||||
@misc{zhang2025fastvideogenerationsliding,
|
||||
title={Fast Video Generation with Sliding Tile Attention},
|
||||
author={Peiyuan Zhang and Yongqi Chen and Runlong Su and Hangliang Ding and Ion Stoica and Zhenghong Liu and Hao Zhang},
|
||||
|
||||
+15
@@ -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'
|
||||
]
|
||||
@@ -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']
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 1.3 MiB |
Binary file not shown.
|
After Width: | Height: | Size: 31 KiB |
Binary file not shown.
|
After Width: | Height: | Size: 8.7 MiB |
Binary file not shown.
|
After Width: | Height: | Size: 769 KiB |
@@ -0,0 +1,645 @@
|
||||
{
|
||||
"id": "23a1f065-bbba-4a8f-b144-944e1318fcbf",
|
||||
"revision": 0,
|
||||
"last_node_id": 8,
|
||||
"last_link_id": 7,
|
||||
"nodes": [
|
||||
{
|
||||
"id": 4,
|
||||
"type": "VAEConfig",
|
||||
"pos": [
|
||||
374.2159423828125,
|
||||
554.85888671875
|
||||
],
|
||||
"size": [
|
||||
334.080078125,
|
||||
322
|
||||
],
|
||||
"flags": {},
|
||||
"order": 0,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "vae_config",
|
||||
"type": "VAE_CONFIG",
|
||||
"links": [
|
||||
2
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VAEConfig"
|
||||
},
|
||||
"widgets_values": [
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"load_encoder": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
},
|
||||
"load_decoder": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
},
|
||||
"tile_sample_min_height": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 256
|
||||
},
|
||||
"tile_sample_min_width": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 256
|
||||
},
|
||||
"tile_sample_min_num_frames": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 16
|
||||
},
|
||||
"tile_sample_stride_height": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 192
|
||||
},
|
||||
"tile_sample_stride_width": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 192
|
||||
},
|
||||
"tile_sample_stride_num_frames": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 12
|
||||
},
|
||||
"blend_num_frames": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 0
|
||||
},
|
||||
"use_tiling": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
},
|
||||
"use_temporal_tiling": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
},
|
||||
"use_parallel_tiling": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 5,
|
||||
"type": "TextEncoderConfig",
|
||||
"pos": [
|
||||
416.4937744140625,
|
||||
953.6171875
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
106
|
||||
],
|
||||
"flags": {},
|
||||
"order": 1,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "text_encoder_config",
|
||||
"type": "TEXT_ENCODER_CONFIG",
|
||||
"links": [
|
||||
7
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "TextEncoderConfig"
|
||||
},
|
||||
"widgets_values": [
|
||||
-99999,
|
||||
-99999,
|
||||
-99999
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"prefix": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": ""
|
||||
},
|
||||
"quant_config": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": ""
|
||||
},
|
||||
"lora_config": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": ""
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 6,
|
||||
"type": "DITConfig",
|
||||
"pos": [
|
||||
415.1928405761719,
|
||||
1154.1573486328125
|
||||
],
|
||||
"size": [
|
||||
270,
|
||||
82
|
||||
],
|
||||
"flags": {},
|
||||
"order": 2,
|
||||
"mode": 0,
|
||||
"inputs": [],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "dit_config",
|
||||
"type": "DIT_CONFIG",
|
||||
"links": [
|
||||
6
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "DITConfig"
|
||||
},
|
||||
"widgets_values": [
|
||||
-99999,
|
||||
-99999
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"prefix": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": ""
|
||||
},
|
||||
"quant_config": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": ""
|
||||
}
|
||||
}
|
||||
},
|
||||
{
|
||||
"id": 1,
|
||||
"type": "VideoGenerator",
|
||||
"pos": [
|
||||
818.804931640625,
|
||||
348.9299621582031
|
||||
],
|
||||
"size": [
|
||||
400,
|
||||
436
|
||||
],
|
||||
"flags": {},
|
||||
"order": 4,
|
||||
"mode": 0,
|
||||
"inputs": [
|
||||
{
|
||||
"name": "inference_args",
|
||||
"shape": 7,
|
||||
"type": "INFERENCE_ARGS",
|
||||
"link": 3
|
||||
},
|
||||
{
|
||||
"name": "vae_config",
|
||||
"shape": 7,
|
||||
"type": "VAE_CONFIG",
|
||||
"link": 2
|
||||
},
|
||||
{
|
||||
"name": "text_encoder_config",
|
||||
"shape": 7,
|
||||
"type": "TEXT_ENCODER_CONFIG",
|
||||
"link": 7
|
||||
},
|
||||
{
|
||||
"name": "dit_config",
|
||||
"shape": 7,
|
||||
"type": "DIT_CONFIG",
|
||||
"link": 6
|
||||
}
|
||||
],
|
||||
"outputs": [
|
||||
{
|
||||
"name": "video_path",
|
||||
"type": "STRING",
|
||||
"links": [
|
||||
4
|
||||
]
|
||||
}
|
||||
],
|
||||
"properties": {
|
||||
"Node name for S&R": "VideoGenerator"
|
||||
},
|
||||
"widgets_values": [
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes wide with interest. The playful yet serene atmosphere is complemented by soft natural light filtering through the petals. Mid-shot, warm and cheerful tones.",
|
||||
"/workspace/ComfyUI/outputs_video/",
|
||||
2,
|
||||
"FastVideo/FastHunyuan-diffusers",
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999,
|
||||
-99999
|
||||
],
|
||||
"auto_widget_states": {
|
||||
"embedded_cfg_scale": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 6
|
||||
},
|
||||
"sp_size": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 2
|
||||
},
|
||||
"tp_size": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": 2
|
||||
},
|
||||
"vae_precision": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": "fp16"
|
||||
},
|
||||
"vae_tiling": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
},
|
||||
"vae_sp": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": true
|
||||
},
|
||||
"text_encoder_precision": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": "fp16"
|
||||
},
|
||||
"precision": {
|
||||
"isAuto": true,
|
||||
"value": -99999,
|
||||
"cachedValue": "fp16"
|
||||
},
|
||||
"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");
|
||||
+46
-16
@@ -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 support H100/H200, because ThunderKittens uses TMA but doesn't support Blackwell yet.
|
||||
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.4)
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
git submodule update --init --recursive
|
||||
```
|
||||
|
||||
## 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,22 +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 tests/test_sta.py # test STA
|
||||
python tests/test_block_sparse.py # test VSA
|
||||
```
|
||||
## Benchmark
|
||||
### 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.
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -0,0 +1,4 @@
|
||||
off_hz = tl.program_id(2)
|
||||
b = off_hz // H
|
||||
h = off_hz % H
|
||||
meta_base = ((b * H + h) * q_tiles + q_blk)
|
||||
+18
-10
@@ -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
|
||||
|
||||
@@ -51,21 +51,29 @@ for k in kernels:
|
||||
source_files.append(sources[k]['source_files'][target])
|
||||
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
|
||||
|
||||
ext_modules = []
|
||||
import torch
|
||||
major, minor = torch.cuda.get_device_capability(0)
|
||||
if major == 9 and minor == 0:# check if H100
|
||||
ext_modules = [
|
||||
CUDAExtension('vsa_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
]
|
||||
|
||||
|
||||
|
||||
setup(name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
ext_modules=[
|
||||
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",
|
||||
|
||||
@@ -1,289 +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
|
||||
import gc
|
||||
|
||||
def set_seed(seed: int = 42):
|
||||
# Python random module
|
||||
random.seed(seed)
|
||||
|
||||
# NumPy
|
||||
np.random.seed(seed)
|
||||
|
||||
# PyTorch
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed) # if using multi-GPU
|
||||
|
||||
|
||||
@torch.no_grad
|
||||
def precision_metric(quant_o, fa2_o):
|
||||
x, xx = quant_o.float(), fa2_o.float()
|
||||
sim = torch.nn.functional.cosine_similarity(x.reshape(1, -1), xx.reshape(1, -1)).item()
|
||||
l1 = ((x - xx).abs().sum() / xx.abs().sum() ).item()
|
||||
rmse = torch.sqrt(torch.mean((x -xx) ** 2)).item()
|
||||
|
||||
return sim, l1, rmse
|
||||
|
||||
def create_input_tensors(batch, head, seq_len, headdim):
|
||||
"""Create random input tensors for attention."""
|
||||
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
return q, k, v
|
||||
|
||||
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
|
||||
|
||||
Args:
|
||||
bs: batch size
|
||||
h: number of heads
|
||||
num_q_blocks: number of query blocks
|
||||
num_kv_blocks: number of key-value blocks
|
||||
k: number of kv blocks each q block attends to
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to k).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Binary mask where 1 indicates attention connection.
|
||||
"""
|
||||
# Ensure k is not larger than num_kv_blocks
|
||||
k = min(k, num_kv_blocks)
|
||||
|
||||
# Create random scores for sampling
|
||||
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
|
||||
|
||||
# Get top-k indices for each q block
|
||||
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
|
||||
|
||||
# sort q2k_block_sparse_index
|
||||
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
|
||||
|
||||
# All q blocks attend to exactly k kv blocks
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
|
||||
|
||||
# Create the corresponding mask
|
||||
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
|
||||
|
||||
# Fill in the mask based on the indices
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx]
|
||||
block_sparse_mask[b, head, q_idx, kv_indices] = True
|
||||
|
||||
# Create the reverse mapping (k2q)
|
||||
# First, initialize lists to collect q indices for each kv block
|
||||
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
|
||||
|
||||
# Populate the lists based on q2k mapping
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
|
||||
for kv_idx in kv_indices:
|
||||
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
|
||||
|
||||
# Find the maximum number of q blocks that attend to any kv block
|
||||
max_q_per_kv = 0
|
||||
for flat_idx in range(bs * h):
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
|
||||
|
||||
# Create tensors for k2q mapping
|
||||
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
|
||||
dtype=torch.int32, device=device)
|
||||
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
|
||||
dtype=torch.int32, device=device)
|
||||
|
||||
# Fill the tensors
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
q_indices = k2q_indices_list[flat_idx][kv_idx]
|
||||
num_q = len(q_indices)
|
||||
k2q_block_sparse_num[b, head, kv_idx] = num_q
|
||||
if num_q > 0:
|
||||
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
|
||||
q_indices, dtype=torch.int32, device=device)
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
|
||||
|
||||
def main(args):
|
||||
set_seed(42)
|
||||
|
||||
# Extract parameters
|
||||
batch = args.batch_size
|
||||
head = args.num_heads
|
||||
headdim = args.head_dim
|
||||
num_iterations = args.num_iterations
|
||||
|
||||
print(f"Block Sparse Attention Benchmark")
|
||||
print(f"batch: {batch}, head: {head}, headdim: {headdim}, iterations: {num_iterations}")
|
||||
|
||||
# Test with different sequence lengths
|
||||
for seq_len in args.seq_lengths:
|
||||
# Skip very long sequences if they might cause OOM
|
||||
# if seq_len > 16384 and batch > 1:
|
||||
# continue
|
||||
|
||||
print("="*100)
|
||||
print(f"\nSequence length: {seq_len}")
|
||||
|
||||
# Collect metrics across iterations
|
||||
forward_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_q_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_k_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_v_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
|
||||
for iter_idx in range(num_iterations):
|
||||
if num_iterations > 1:
|
||||
print(f"\nIteration {iter_idx+1}/{num_iterations}")
|
||||
|
||||
# Create input tensors
|
||||
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
|
||||
|
||||
# Setup block sparse parameters
|
||||
num_q_blocks = seq_len // BLOCK_M
|
||||
num_kv_blocks = seq_len // BLOCK_N
|
||||
|
||||
# Determine k value (number of kv blocks per q block)
|
||||
topk = args.topk
|
||||
if topk is None:
|
||||
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
|
||||
topk = max(1, topk)
|
||||
if iter_idx == 0: # Only print this once
|
||||
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
|
||||
|
||||
# Generate block sparse pattern
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask = generate_block_sparse_pattern(
|
||||
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
|
||||
|
||||
# expand block_sparse_mask to full mask
|
||||
block_mask_expanded = block_sparse_mask.unsqueeze(-1).unsqueeze(-2) # [b, h, num_q_blocks, num_kv_blocks, 1, 1]
|
||||
block_mask_expanded = block_mask_expanded.expand(-1, -1, -1, -1, BLOCK_M, BLOCK_N) # [b, h, num_q_blocks, num_kv_blocks, BLOCK_M, BLOCK_N]
|
||||
full_mask = block_mask_expanded.permute(0, 1, 2, 4, 3, 5).reshape(batch, head, seq_len, seq_len)
|
||||
|
||||
q.requires_grad = True
|
||||
k.requires_grad = True
|
||||
v.requires_grad = True
|
||||
|
||||
|
||||
# testing forward
|
||||
o = BlockSparseAttentionFunction.apply(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
del q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask, block_mask_expanded
|
||||
grad_o = torch.randn_like(o)
|
||||
o.backward(grad_o)
|
||||
# clear memory
|
||||
q_sdpa = q.detach().clone()
|
||||
k_sdpa = k.detach().clone()
|
||||
v_sdpa = v.detach().clone()
|
||||
q_sdpa.requires_grad = True
|
||||
k_sdpa.requires_grad = True
|
||||
v_sdpa.requires_grad = True
|
||||
q.data = torch.empty(0, device=q.device)
|
||||
k.data = torch.empty(0, device=k.device)
|
||||
v.data = torch.empty(0, device=v.device)
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
o_sdpa = torch.nn.functional.scaled_dot_product_attention(q_sdpa, k_sdpa, v_sdpa, attn_mask=full_mask)
|
||||
|
||||
|
||||
sim, l1, rmse = precision_metric(o, o_sdpa)
|
||||
assert sim > 0.9999, f"SSIM too low: {sim}"
|
||||
assert l1 < 8e-5, f"l1 too large: {l1}"
|
||||
assert rmse < 2e-5, f"RMSE too large: {rmse}"
|
||||
forward_metrics['sim'].append(sim)
|
||||
forward_metrics['l1'].append(l1)
|
||||
forward_metrics['rmse'].append(rmse)
|
||||
|
||||
print(f"block_sparse_attention_fwd vs torch.nn.functional.scaled_dot_product_attention:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
# test backward
|
||||
o_sdpa.backward(grad_o)
|
||||
|
||||
sim, l1, rmse = precision_metric(q.grad, q_sdpa.grad)
|
||||
# Error bounds collected on H100
|
||||
assert sim > 0.9999, f"SSIM too low: {sim}"
|
||||
assert l1 < 4e-3, f"l1 too large: {l1}"
|
||||
assert rmse < 3e-4, f"RMSE too large: {rmse}"
|
||||
grad_q_metrics['sim'].append(sim)
|
||||
grad_q_metrics['l1'].append(l1)
|
||||
grad_q_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_q:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
sim, l1, rmse = precision_metric(k.grad, k_sdpa.grad)
|
||||
assert sim > 0.9999, f"SSIM too low: {sim}"
|
||||
assert l1 < 4e-3, f"l1 too large: {l1}"
|
||||
assert rmse < 2e-4, f"RMSE too large: {rmse}"
|
||||
grad_k_metrics['sim'].append(sim)
|
||||
grad_k_metrics['l1'].append(l1)
|
||||
grad_k_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_k:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
sim, l1, rmse = precision_metric(v.grad, v_sdpa.grad)
|
||||
assert sim > 0.9999, f"SSIM too low: {sim}"
|
||||
assert l1 < 1e-4, f"l1 too large: {l1}"
|
||||
assert rmse < 2e-5, f"RMSE too large: {rmse}"
|
||||
grad_v_metrics['sim'].append(sim)
|
||||
grad_v_metrics['l1'].append(l1)
|
||||
grad_v_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_v:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
del o, o_sdpa, grad_o, q_sdpa, k_sdpa, v_sdpa
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Print summary statistics if multiple iterations were run
|
||||
if num_iterations > 1:
|
||||
print("\n" + "="*50)
|
||||
print(f"Summary Statistics (over {num_iterations} iterations):")
|
||||
|
||||
print("\nForward metrics:")
|
||||
print(f"Similarity: mean={np.mean(forward_metrics['sim']):.6f}, std={np.std(forward_metrics['sim']):.6f}, min={np.min(forward_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(forward_metrics['l1']):.6f}, std={np.std(forward_metrics['l1']):.6f}, max={np.max(forward_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(forward_metrics['rmse']):.6f}, std={np.std(forward_metrics['rmse']):.6f}, max={np.max(forward_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient Q metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_q_metrics['sim']):.6f}, std={np.std(grad_q_metrics['sim']):.6f}, min={np.min(grad_q_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_q_metrics['l1']):.6f}, std={np.std(grad_q_metrics['l1']):.6f}, max={np.max(grad_q_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_q_metrics['rmse']):.6f}, std={np.std(grad_q_metrics['rmse']):.6f}, max={np.max(grad_q_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient K metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_k_metrics['sim']):.6f}, std={np.std(grad_k_metrics['sim']):.6f}, min={np.min(grad_k_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_k_metrics['l1']):.6f}, std={np.std(grad_k_metrics['l1']):.6f}, max={np.max(grad_k_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_k_metrics['rmse']):.6f}, std={np.std(grad_k_metrics['rmse']):.6f}, max={np.max(grad_k_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient V metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_v_metrics['sim']):.6f}, std={np.std(grad_v_metrics['sim']):.6f}, min={np.min(grad_v_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_v_metrics['l1']):.6f}, std={np.std(grad_v_metrics['l1']):.6f}, max={np.max(grad_v_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_v_metrics['rmse']):.6f}, std={np.std(grad_v_metrics['rmse']):.6f}, max={np.max(grad_v_metrics['rmse']):.6f}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
|
||||
parser.add_argument('--batch_size', type=int, default=4, help='Batch size')
|
||||
parser.add_argument('--num_heads', type=int, default=6, help='Number of heads')
|
||||
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=64, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[29120], help='Sequence lengths to benchmark')
|
||||
parser.add_argument('--num_iterations', type=int, default=50, help='Number of test iterations to run')
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -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.")
|
||||
@@ -0,0 +1,159 @@
|
||||
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().requires_grad_()
|
||||
k_ = K.clone().requires_grad_()
|
||||
v_ = V.clone().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.clone().requires_grad_()
|
||||
K = K.clone().requires_grad_()
|
||||
V = V.clone().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, 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 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, mean, std, 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), mean, std, torch.bfloat16, device)
|
||||
K = generate_tensor((1, h, S, d), mean, std, torch.bfloat16, device)
|
||||
V = generate_tensor((1, h, S, d), mean, std, torch.bfloat16, device)
|
||||
dO = generate_tensor((1, h, S, d), mean, std, 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()
|
||||
results[name]['max_diff'] = max(results[name]['max_diff'], torch.max(abs_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, mean, std, 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}, mean={mean}, std={std}, mode={error_mode}")
|
||||
print("=" * 150)
|
||||
print(f"{'Config':<20} {'Blocks':<8} {'K':<4} "
|
||||
f"{'gQ Avg':<12} {'gQ Max':<12} "
|
||||
f"{'gK Avg':<12} {'gK Max':<12} "
|
||||
f"{'gV Avg':<12} {'gV Max':<12} "
|
||||
f"{'gO Avg':<12} {'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, mean, std, 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
|
||||
mean = 0.0
|
||||
std = 1
|
||||
print("Block Sparse Attention with Variable Block Sizes Analysis")
|
||||
print("=" * 60)
|
||||
for mode in ['backward']:
|
||||
generate_error_graphs(h, d, mean, std, 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
|
||||
+2
-2
@@ -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
|
||||
|
||||
|
||||
+47
-438
@@ -1,71 +1,20 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
from torch.utils.checkpoint import detach_variable
|
||||
from typing import Tuple
|
||||
block_sparse_attn=None
|
||||
|
||||
try:
|
||||
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
|
||||
except ImportError:
|
||||
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 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))
|
||||
@@ -77,394 +26,54 @@ def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
output = torch.matmul(QK, v)
|
||||
return output, QK
|
||||
|
||||
def generate_topk_block_sparse_pattern(block_attn_score: torch.Tensor,
|
||||
topk: int):
|
||||
|
||||
def video_sparse_attn(q, k, v, variable_block_sizes, topk, block_size, compress_attn_weight=None):
|
||||
"""
|
||||
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.
|
||||
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]]
|
||||
"""
|
||||
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)
|
||||
if isinstance(block_size, int):
|
||||
block_size = (block_size, block_size, block_size)
|
||||
|
||||
sorted_indice = sorted_result.indices
|
||||
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)
|
||||
|
||||
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)
|
||||
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
|
||||
v_compress)
|
||||
|
||||
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))
|
||||
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)
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
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)
|
||||
|
||||
@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)
|
||||
if compress_attn_weight is not None:
|
||||
final_output = output_compress * compress_attn_weight + output_select
|
||||
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)
|
||||
final_output = output_compress + output_select
|
||||
return final_output
|
||||
|
||||
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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -46,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)
|
||||
@@ -357,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) {
|
||||
@@ -409,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) {
|
||||
@@ -484,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]);
|
||||
@@ -496,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;
|
||||
@@ -507,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);
|
||||
}
|
||||
@@ -528,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);
|
||||
}
|
||||
}
|
||||
@@ -591,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
|
||||
@@ -713,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;
|
||||
@@ -732,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);
|
||||
@@ -810,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);
|
||||
@@ -876,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);
|
||||
@@ -888,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");
|
||||
@@ -967,7 +768,19 @@ 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())
|
||||
};
|
||||
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
@@ -1006,7 +819,19 @@ 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())
|
||||
};
|
||||
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
@@ -1036,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);
|
||||
@@ -1049,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");
|
||||
@@ -1136,7 +962,7 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
float* d_vg = reinterpret_cast<float*>(vg_ptr);
|
||||
|
||||
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
|
||||
int threads = 4 * kittens::WARP_THREADS;
|
||||
int threads = PREP_NUM_WARPS * kittens::WARP_THREADS;
|
||||
|
||||
//cudadevicesynchronize();
|
||||
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
|
||||
@@ -1145,7 +971,7 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
// 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>;
|
||||
@@ -1220,8 +1046,8 @@ 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;
|
||||
@@ -1324,8 +1150,8 @@ 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;
|
||||
|
||||
@@ -0,0 +1,182 @@
|
||||
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_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)
|
||||
|
||||
|
||||
@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)
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
```
|
||||
|
||||
+1
-3
@@ -18,7 +18,6 @@ import os
|
||||
import re
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
import requests
|
||||
|
||||
@@ -168,8 +167,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
|
||||
|
||||
@@ -8,12 +8,14 @@ You can easily use the FastVideo Docker image as a custom container on [RunPod](
|
||||
|
||||
Choose a GPU that supports CUDA 12.4
|
||||
|
||||
Pick 1 or 2 L40S GPU(s)
|
||||
|
||||

|
||||
|
||||
When creating your pod template, use this image:
|
||||
|
||||
```
|
||||
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:latest
|
||||
ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:py3.12-latest
|
||||
```
|
||||
|
||||
Paste Container Start Command to support SSH ([RunPod Docs](https://docs.runpod.io/pods/configuration/use-ssh)):
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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,173 @@ def generate_examples(generate_main_index=False):
|
||||
continue
|
||||
examples.append(Example(path.parent))
|
||||
|
||||
# Create document directories for each category based on category name and generate files
|
||||
for example in sorted(examples, key=lambda e: e.path.stem):
|
||||
print(example)
|
||||
return examples
|
||||
|
||||
|
||||
def create_nested_structures(
|
||||
examples: list[Example]
|
||||
) -> dict[str, dict[str, dict[str, dict[str, NestedStructure]]]]:
|
||||
"""Create nested structures for training and distillation categories."""
|
||||
nested_structures: dict[str, dict[str, dict[str,
|
||||
dict[str,
|
||||
NestedStructure]]]] = {}
|
||||
|
||||
for example in examples:
|
||||
if example.category not in ["training", "distillation"]:
|
||||
continue
|
||||
|
||||
category_dir = EXAMPLE_DIR / example.category
|
||||
relative_path = example.path.relative_to(category_dir)
|
||||
path_parts = relative_path.parts
|
||||
|
||||
# For nested examples like finetune/wan_i2v_14b_480p/crush_smol
|
||||
if len(path_parts) >= 3:
|
||||
method = path_parts[0] # e.g., "finetune"
|
||||
model = path_parts[1] # e.g., "wan_i2v_14b_480p"
|
||||
dataset = path_parts[2] # e.g., "crush_smol"
|
||||
|
||||
# Initialize nested structure
|
||||
if example.category not in nested_structures:
|
||||
nested_structures[example.category] = {}
|
||||
if method not in nested_structures[example.category]:
|
||||
nested_structures[example.category][method] = {}
|
||||
if model not in nested_structures[example.category][method]:
|
||||
nested_structures[example.category][method][model] = {}
|
||||
|
||||
# Store the nested structure
|
||||
nested_structures[
|
||||
example.category][method][model][dataset] = NestedStructure(
|
||||
category=example.category,
|
||||
method=method,
|
||||
model=model,
|
||||
dataset=dataset,
|
||||
example=example)
|
||||
|
||||
return nested_structures
|
||||
|
||||
|
||||
def generate_flat_examples(examples: list[Example],
|
||||
category_indices: dict[str, Index],
|
||||
examples_index: Index | None,
|
||||
generate_main_index: bool) -> None:
|
||||
"""Generate documentation for flat structure examples (inference, etc.)."""
|
||||
for example in examples:
|
||||
if example.category in ["training", "distillation"]:
|
||||
continue # Skip nested structure examples
|
||||
|
||||
# Determine which index to use for this example
|
||||
if example.category is not None and example.category in category_indices:
|
||||
index = category_indices[example.category]
|
||||
elif generate_main_index:
|
||||
assert examples_index is not None
|
||||
index = examples_index # Default to main index if available
|
||||
index = examples_index
|
||||
else:
|
||||
# Skip examples without a category if no main index
|
||||
print(f"Skipping {example.path} (no category and no main index)")
|
||||
continue
|
||||
|
||||
# Place generated example markdown in the same directory as its index
|
||||
# Generate the example documentation
|
||||
doc_path = index.path.parent / f"{example.path.stem}.md"
|
||||
with open(doc_path, "w+") as f:
|
||||
f.write(example.generate())
|
||||
# Add the example to the index
|
||||
index.documents.append(example.path.stem)
|
||||
|
||||
|
||||
def generate_nested_examples(nested_structures: dict[str, dict[str, dict[
|
||||
str, dict[str, NestedStructure]]]], category_indices: dict[str,
|
||||
Index]) -> None:
|
||||
"""Generate documentation for nested structure examples (training, distillation)."""
|
||||
for category_name in ["training", "distillation"]:
|
||||
if category_name not in category_indices or category_name not in nested_structures:
|
||||
continue
|
||||
|
||||
category_index = category_indices[category_name]
|
||||
category_base_dir = category_index.path.parent
|
||||
|
||||
for method, models in nested_structures[category_name].items():
|
||||
# Create method-level index
|
||||
method_index = Index(path=category_base_dir / f"{method}.md",
|
||||
title=fix_case(method),
|
||||
description=f"Examples using {method}.",
|
||||
caption=f"{fix_case(method)} Examples",
|
||||
maxdepth=2)
|
||||
|
||||
for model, datasets in models.items():
|
||||
# Generate dataset examples using the Example class
|
||||
for dataset, nested_struct in datasets.items():
|
||||
doc_path = category_base_dir / f"{nested_struct.filename}.md"
|
||||
with open(doc_path, "w+") as f:
|
||||
f.write(nested_struct.example.generate())
|
||||
|
||||
# Create model-level index
|
||||
model_index = Index(
|
||||
path=category_base_dir / f"{model}.md",
|
||||
title=fix_case(model.replace('_', ' ')),
|
||||
description=f"Examples for the {model} model.",
|
||||
caption=f"{fix_case(model.replace('_', ' '))} Datasets",
|
||||
maxdepth=1)
|
||||
|
||||
# Add dataset indices to model index
|
||||
for dataset, nested_struct in datasets.items():
|
||||
model_index.documents.append(nested_struct.filename)
|
||||
|
||||
# Write model index
|
||||
with open(model_index.path, "w+") as f:
|
||||
f.write(model_index.generate())
|
||||
|
||||
# Add model to method index
|
||||
method_index.documents.append(model)
|
||||
|
||||
# Write method index
|
||||
with open(method_index.path, "w+") as f:
|
||||
f.write(method_index.generate())
|
||||
|
||||
# Add method to main category index
|
||||
category_index.documents.append(method)
|
||||
|
||||
|
||||
def generate_examples(generate_main_index=False):
|
||||
"""
|
||||
Generate example documentation.
|
||||
|
||||
Args:
|
||||
generate_main_index (bool): Whether to generate the main examples index.
|
||||
If False, only category-specific indices will be generated.
|
||||
"""
|
||||
# Create category indices
|
||||
category_indices = create_category_indices()
|
||||
|
||||
# Create the main examples index only if requested
|
||||
examples_index = None
|
||||
if generate_main_index:
|
||||
main_index_dir = ROOT_DIR / "docs/source/examples"
|
||||
examples_index = Index(
|
||||
path=main_index_dir / "examples_index.md",
|
||||
title="💡 Examples",
|
||||
description=
|
||||
"A collection of examples demonstrating usage of FastVideo.\nAll documented examples are autogenerated using <gh-file:docs/source/generate_examples.py> from examples found in <gh-file:examples>.",
|
||||
caption="Examples",
|
||||
maxdepth=2)
|
||||
|
||||
# Find all examples
|
||||
examples = find_examples(category_indices, generate_main_index)
|
||||
|
||||
# Create nested structures for training and distillation
|
||||
nested_structures = create_nested_structures(examples)
|
||||
|
||||
# Generate flat structure examples (inference, etc.)
|
||||
generate_flat_examples(examples, category_indices, examples_index,
|
||||
generate_main_index)
|
||||
|
||||
# Generate nested structure examples (training, distillation)
|
||||
generate_nested_examples(nested_structures, category_indices)
|
||||
|
||||
# Generate the index files for categories
|
||||
for category_index in category_indices.values():
|
||||
if category_index.documents:
|
||||
# Add to main index if it exists
|
||||
if generate_main_index:
|
||||
if generate_main_index and examples_index:
|
||||
main_index_dir = examples_index.path.parent
|
||||
rel_path = category_index.path.relative_to(
|
||||
main_index_dir.parent)
|
||||
assert examples_index is not None
|
||||
examples_index.documents.insert(
|
||||
0,
|
||||
str(rel_path).replace(".md", ""))
|
||||
|
||||
@@ -1,120 +1,18 @@
|
||||
(fastvideo-installation)=
|
||||
(installation-index)=
|
||||
|
||||
# 🔧 Installation
|
||||
|
||||
FastVideo currently only supports Linux and NVIDIA CUDA GPUs.
|
||||
FastVideo supports the following hardware platforms:
|
||||
|
||||
## Requirements
|
||||
:::{toctree}
|
||||
:maxdepth: 1
|
||||
:hidden:
|
||||
|
||||
- **OS: Linux**
|
||||
- **Python: 3.10-3.12**
|
||||
- **CUDA 12.4**
|
||||
- **At least 1 NVIDIA GPU**
|
||||
|
||||
## Set up using Python
|
||||
### Create a new Python environment
|
||||
|
||||
#### Conda
|
||||
You can create a new python environment using [Conda](https://docs.conda.io/projects/conda/en/stable/user-guide/getting-started.html)
|
||||
##### 1. Install Miniconda (if not already installed)
|
||||
|
||||
```bash
|
||||
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
|
||||
bash Miniconda3-latest-Linux-x86_64.sh
|
||||
source ~/.bashrc
|
||||
```
|
||||
|
||||
##### 2. Create and activate a Conda environment for FastVideo
|
||||
|
||||
```bash
|
||||
# (Recommended) Create a new conda environment.
|
||||
conda create -n fastvideo python=3.12 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
:::{note}
|
||||
[PyTorch has deprecated the conda release channel](https://github.com/pytorch/pytorch/issues/138506). If you use `conda`, please only use it to create Python environment rather than installing packages.
|
||||
installation/gpu
|
||||
installation/mps
|
||||
:::
|
||||
|
||||
#### uv
|
||||
|
||||
:::{tip}
|
||||
We highly recommend using `uv` to install FastVideo. In our experience, `uv` speeds up installation by at least 3x.
|
||||
:::
|
||||
|
||||
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
|
||||
|
||||
```console
|
||||
# (Recommended) Create a new uv environment. Use `--seed` to install `pip` and `setuptools` in the environment.
|
||||
uv venv --python 3.12 --seed
|
||||
source .venv/bin/activate
|
||||
```
|
||||
|
||||
### Installation
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
|
||||
# or if you are using uv
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
Also optionally install flash-attn:
|
||||
|
||||
```bash
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
|
||||
#### 1. Clone the FastVideo repository
|
||||
|
||||
```bash
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
```
|
||||
|
||||
#### 2. Install FastVideo
|
||||
|
||||
Basic installation:
|
||||
|
||||
```bash
|
||||
pip install -e .
|
||||
|
||||
# or if you are using uv
|
||||
uv pip install -e .
|
||||
```
|
||||
|
||||
### Optional Dependencies
|
||||
|
||||
#### Flash Attention
|
||||
|
||||
```bash
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
## Set up using Docker
|
||||
We also have prebuilt docker images with FastVideo dependencies pre-installed:
|
||||
[Docker Images](#docker)
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
If you're planning to contribute to FastVideo please see the following page:
|
||||
[Contributor Guide](#developer-overview)
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
### For Basic Inference
|
||||
- NVIDIA GPU with CUDA 12.4 support
|
||||
|
||||
### For Lora Finetuning
|
||||
- 40GB GPU memory each for 2 GPUs with lora
|
||||
- 30GB GPU memory each for 2 GPUs with CPU offload and lora
|
||||
|
||||
### For Full Finetuning/Distillation
|
||||
- Multiple high-memory GPUs recommended (e.g., H100)
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
If you encounter any issues during installation, please open an issue on our [GitHub repository](https://github.com/hao-ai-lab/FastVideo).
|
||||
|
||||
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg) for additional support.
|
||||
- <project:installation/gpu.md>
|
||||
- NVIDIA CUDA
|
||||
- <project:installation/mps.md>
|
||||
- Apple silicon
|
||||
|
||||
@@ -0,0 +1,118 @@
|
||||
# NVIDIA GPU
|
||||
|
||||
Instructions to install FastVideo for NVIDIA CUDA GPUs.
|
||||
|
||||
## Requirements
|
||||
|
||||
- **OS: Linux or Windows WSL**
|
||||
- **Python: 3.10-3.12**
|
||||
- **CUDA 12.4**
|
||||
- **At least 1 NVIDIA GPU**
|
||||
|
||||
## Set up using Python
|
||||
### Create a new Python environment
|
||||
|
||||
#### Conda
|
||||
You can create a new python environment using [Conda](https://docs.conda.io/projects/conda/en/stable/user-guide/getting-started.html)
|
||||
##### 1. Install Miniconda (if not already installed)
|
||||
|
||||
```bash
|
||||
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
|
||||
bash Miniconda3-latest-Linux-x86_64.sh
|
||||
source ~/.bashrc
|
||||
```
|
||||
|
||||
##### 2. Create and activate a Conda environment for FastVideo
|
||||
|
||||
```bash
|
||||
# (Recommended) Create a new conda environment.
|
||||
conda create -n fastvideo python=3.12 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
:::{note}
|
||||
[PyTorch has deprecated the conda release channel](https://github.com/pytorch/pytorch/issues/138506). If you use `conda`, please only use it to create Python environment rather than installing packages.
|
||||
:::
|
||||
|
||||
#### uv
|
||||
|
||||
:::{tip}
|
||||
We highly recommend using `uv` to install FastVideo. In our experience, `uv` speeds up installation by at least 3x.
|
||||
:::
|
||||
|
||||
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
|
||||
|
||||
```console
|
||||
# (Recommended) Create a new uv environment. Use `--seed` to install `pip` and `setuptools` in the environment.
|
||||
uv venv --python 3.12 --seed
|
||||
source .venv/bin/activate
|
||||
```
|
||||
|
||||
### Installation
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
|
||||
# or if you are using uv
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
Also optionally install flash-attn:
|
||||
|
||||
```bash
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
|
||||
#### 1. Clone the FastVideo repository
|
||||
|
||||
```bash
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
```
|
||||
|
||||
#### 2. Install FastVideo
|
||||
|
||||
Basic installation:
|
||||
|
||||
```bash
|
||||
pip install -e .
|
||||
|
||||
# or if you are using uv
|
||||
uv pip install -e .
|
||||
```
|
||||
|
||||
### Optional Dependencies
|
||||
|
||||
#### Flash Attention
|
||||
|
||||
```bash
|
||||
pip install flash-attn==2.7.4.post1 --no-build-isolation
|
||||
```
|
||||
|
||||
## Set up using Docker
|
||||
We also have prebuilt docker images with FastVideo dependencies pre-installed:
|
||||
[Docker Images](#docker)
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
If you're planning to contribute to FastVideo please see the following page:
|
||||
[Contributor Guide](#developer-overview)
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
### For Basic Inference
|
||||
- NVIDIA GPU with CUDA 12.4 support
|
||||
|
||||
### For Lora Finetuning
|
||||
- 40GB GPU memory each for 2 GPUs with lora
|
||||
- 30GB GPU memory each for 2 GPUs with CPU offload and lora
|
||||
|
||||
### For Full Finetuning/Distillation
|
||||
- Multiple high-memory GPUs recommended (e.g., H100)
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
If you encounter any issues during installation, please open an issue on our [GitHub repository](https://github.com/hao-ai-lab/FastVideo).
|
||||
|
||||
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ) for additional support.
|
||||
@@ -0,0 +1,101 @@
|
||||
# MPS (Apple Silicon)
|
||||
|
||||
Instructions to install FastVideo for Apple Silicon.
|
||||
|
||||
## Requirements
|
||||
|
||||
- **OS: MacOS**
|
||||
- **Python: 3.12.4**
|
||||
|
||||
## Set up using Python
|
||||
|
||||
### Create a new Python environment
|
||||
|
||||
#### Conda
|
||||
|
||||
You can create a new python environment using [Conda](https://docs.conda.io/projects/conda/en/stable/user-guide/getting-started.html)
|
||||
|
||||
##### 1. Install Miniconda (if not already installed)
|
||||
|
||||
```bash
|
||||
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-MacOSX-arm64.sh
|
||||
bash Miniconda3-latest-MacOSX-arm64.sh
|
||||
source ~/.zshrc
|
||||
```
|
||||
|
||||
##### 2. Create and activate a Conda environment for FastVideo
|
||||
|
||||
```bash
|
||||
# (Recommended) Create a new conda environment.
|
||||
conda create -n fastvideo python=3.12.4 -y
|
||||
conda activate fastvideo
|
||||
```
|
||||
|
||||
:::{note}
|
||||
[PyTorch has deprecated the conda release channel](https://github.com/pytorch/pytorch/issues/138506). If you use `conda`, please only use it to create Python environment rather than installing packages.
|
||||
:::
|
||||
|
||||
#### uv
|
||||
|
||||
:::{tip}
|
||||
We highly recommend using `uv` to install FastVideo. In our experience, `uv` speeds up installation by at least 3x.
|
||||
:::
|
||||
|
||||
Or you can create a new Python environment using [uv](https://docs.astral.sh/uv/), a very fast Python environment manager. Please follow the [documentation](https://docs.astral.sh/uv/#getting-started) to install `uv`. After installing `uv`, you can create a new Python environment using the following command:
|
||||
|
||||
```console
|
||||
# (Recommended) Create a new uv environment. Use `--seed` to install `pip` and `setuptools` in the environment.
|
||||
uv venv --python 3.12 --seed
|
||||
source .venv/bin/activate
|
||||
```
|
||||
|
||||
### Dependencies
|
||||
|
||||
```
|
||||
brew install ffmpeg
|
||||
```
|
||||
|
||||
### Installation
|
||||
|
||||
```bash
|
||||
pip install fastvideo
|
||||
|
||||
# or if you are using uv
|
||||
uv pip install fastvideo
|
||||
```
|
||||
|
||||
### Installation from Source
|
||||
|
||||
#### 1. Clone the FastVideo repository
|
||||
|
||||
```bash
|
||||
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
```
|
||||
|
||||
#### 2. Install FastVideo
|
||||
|
||||
Basic installation:
|
||||
|
||||
```bash
|
||||
pip install -e .
|
||||
|
||||
# or if you are using uv
|
||||
uv pip install -e .
|
||||
```
|
||||
|
||||
## Development Environment Setup
|
||||
|
||||
If you're planning to contribute to FastVideo please see the following page:
|
||||
[Contributor Guide](#developer-overview)
|
||||
|
||||
## Hardware Requirements
|
||||
|
||||
### For Basic Inference
|
||||
|
||||
- Mac M1, M2, M3, or M4 (at least 32 GB RAM is preferable for high quality video generation)
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
If you encounter any issues during installation, please open an issue on our [GitHub repository](https://github.com/hao-ai-lab/FastVideo).
|
||||
|
||||
You can also join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ) for additional support.
|
||||
@@ -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
|
||||
```
|
||||
|
||||
+14
-3
@@ -63,22 +63,33 @@ getting_started/installation
|
||||
:maxdepth: 1
|
||||
|
||||
inference/inference_quick_start
|
||||
inference/examples/examples_inference_index
|
||||
inference/configuration
|
||||
inference/optimizations
|
||||
inference/comfyui
|
||||
inference/support_matrix
|
||||
inference/examples/examples_inference_index
|
||||
inference/cli
|
||||
inference/add_pipeline
|
||||
inference/v0_inference
|
||||
:::
|
||||
|
||||
:::{toctree}
|
||||
:caption: Training
|
||||
:maxdepth: 1
|
||||
|
||||
training/examples/examples_training_index
|
||||
training/data_preprocess
|
||||
training/distillation
|
||||
training/finetune
|
||||
<!-- training/finetune -->
|
||||
:::
|
||||
|
||||
<!-- :::{toctree}
|
||||
:caption: Distillation
|
||||
:maxdepth: 1
|
||||
|
||||
distillation/examples/examples_distillation_index
|
||||
distillation/data_preprocess
|
||||
distillation/dmd -->
|
||||
<!-- training/finetune -->
|
||||
:::
|
||||
|
||||
% What is STA Kernel?
|
||||
|
||||
@@ -12,7 +12,7 @@ This guide explains how to implement a custom diffusion pipeline in FastVideo, l
|
||||
4. **Register Your Pipeline** - Make it discoverable by the framework
|
||||
5. **Configure Your Pipeline** - (Coming soon)
|
||||
|
||||
Need help? Join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg).
|
||||
Need help? Join our [Slack community](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).
|
||||
|
||||
## Step 1: Pipeline Modules
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -121,4 +121,4 @@ If the generated video doesn't match your prompt:
|
||||
- Learn about using [Optimizations](#inference-optimizations)
|
||||
- See [Examples](../examples/examples_inference_index.md) for more usage scenarios
|
||||
- Join our [Community Discord](https://discord.gg/JA7cksDz86).
|
||||
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-2zf6ru791-sRwI9lPIUJQq1mIeB_yjJg).
|
||||
- Join our [Community Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ).
|
||||
|
||||
@@ -1,74 +0,0 @@
|
||||
(v0-inference)=
|
||||
|
||||
# [Deprecated] V0 Inference
|
||||
The following commands and APIs are deprecated but still supported until V1's API can completely replace all the features in this page.
|
||||
|
||||
## Inference StepVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
```
|
||||
python scripts/huggingface/download_hf.py --repo_id=stepfun-ai/stepvideo-t2v --local_dir=data/stepvideo-t2v --repo_type=model
|
||||
```
|
||||
|
||||
Use the following scripts to run inference for StepVideo. When using STA for inference, the generated videos will have dimensions of 204×768×768 (currently, this is the only supported shape).
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_stepvideo_STA.sh # Inference stepvideo with STA
|
||||
sh scripts/inference/inference_stepvideo.sh # Inference original stepvideo
|
||||
```
|
||||
|
||||
## Inference HunyuanVideo with Sliding Tile Attention
|
||||
First, download the model:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/hunyuan --local_dir=data/hunyuan --repo_type=model
|
||||
```
|
||||
|
||||
We provide two examples in the following script to run inference with STA + [TeaCache](https://github.com/ali-vilab/TeaCache) and STA only.
|
||||
|
||||
```bash
|
||||
sh scripts/inference/inference_hunyuan_STA.sh
|
||||
```
|
||||
|
||||
## Video Demos using STA + Teacache
|
||||
Visit our [demo website](https://fast-video.github.io/) to explore our complete collection of examples. We shorten a single video generation process from 945s to 317s on H100.
|
||||
|
||||
## Inference FastHunyuan on single RTX4090
|
||||
We now support NF4 and LLM-INT8 quantized inference using BitsAndBytes for FastHunyuan. With NF4 quantization, inference can be performed on a single RTX 4090 GPU, requiring just 20GB of VRAM.
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan-diffusers --local_dir=data/FastHunyuan-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan_hf_quantization.sh
|
||||
```
|
||||
|
||||
For more information about the VRAM requirements for BitsAndBytes quantization, please refer to the table below (timing measured on an H100 GPU):
|
||||
|
||||
| Configuration | Memory to Init Transformer | Peak Memory After Init Pipeline (Denoise) | Diffusion Time | End-to-End Time |
|
||||
|--------------------------------|----------------------------|--------------------------------------------|----------------|-----------------|
|
||||
| BF16 + Pipeline CPU Offload | 23.883G | 33.744G | 81s | 121.5s |
|
||||
| INT8 + Pipeline CPU Offload | 13.911G | 27.979G | 88s | 116.7s |
|
||||
| NF4 + Pipeline CPU Offload | 9.453G | 19.26G | 78s | 114.5s |
|
||||
|
||||
For improved quality in generated videos, we recommend using a GPU with 80GB of memory to run the BF16 model with the original Hunyuan pipeline. To execute the inference, use the following section:
|
||||
|
||||
## FastHunyuan
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastHunyuan --local_dir=data/FastHunyuan --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_hunyuan.sh
|
||||
```
|
||||
|
||||
You can also inference FastHunyuan in the [official Hunyuan github](https://github.com/Tencent/HunyuanVideo).
|
||||
|
||||
## FastMochi
|
||||
|
||||
```bash
|
||||
# Download the model weight
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/FastMochi-diffusers --local_dir=data/FastMochi-diffusers --repo_type=model
|
||||
# CLI inference
|
||||
bash scripts/inference/inference_mochi_sp.sh
|
||||
```
|
||||
@@ -7,7 +7,7 @@ To save GPU memory, we precompute text embeddings and VAE latents to eliminate t
|
||||
We provide a sample dataset to help you get started. Download the source media using the following command:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/mini_i2v_dataset --local_dir=FastVideo/mini_i2v_dataset --repo_type=dataset
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/mini_i2v_dataset --local_dir=data/mini_i2v_dataset --repo_type=dataset
|
||||
```
|
||||
|
||||
The folder `crush-smol_raw/` contains raw videos and captions for testing preprocessing, while `crush-smol_preprocessed/` contains latents prepared for testing training.
|
||||
|
||||
@@ -0,0 +1,27 @@
|
||||
# Wan2.1-T2V-1.3B Distill Example
|
||||
These are end-to-end example scripts for distilling Wan2.1 T2V 1.3B model using DMD-only and DMD+VSA methods.
|
||||
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
cd csrc/attn
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
### 1. Download dataset:
|
||||
```bash
|
||||
bash examples/distill/Wan-Syn-480P/download_dataset.sh
|
||||
```
|
||||
|
||||
### 2. Configure and run distillation:
|
||||
|
||||
#### For DMD-only distillation:
|
||||
```bash
|
||||
sbatch examples/distill/Wan-Syn-480P/distill_dmd_t2v.slurm
|
||||
```
|
||||
|
||||
#### For DMD+VSA distillation:
|
||||
```bash
|
||||
sbatch examples/distill/Wan-Syn-480P/distill_dmd_VSA_t2v.slurm
|
||||
```
|
||||
@@ -0,0 +1,138 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=8
|
||||
#SBATCH --ntasks=8
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=dmd_t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate your_env
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_distill_dmd_VSA
|
||||
--output_dir"checkpoints/wan_t2v_finetune"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 64
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "3"
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--ema_start_step 0
|
||||
--flow_shift 8
|
||||
--seed 1000
|
||||
)
|
||||
|
||||
# DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,757,522'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--generator_update_interval 5
|
||||
--real_score_guidance_scale 3.5
|
||||
--VSA_sparsity 0.8
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/wan_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
@@ -0,0 +1,138 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=8
|
||||
#SBATCH --ntasks=8
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=dmd_t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate your_env
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_distill_dmd_VSA
|
||||
--output_dir "checkpoints/wan_t2v_finetune"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus 64
|
||||
--sp_size 4
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 8
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "3"
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--ema_start_step 0
|
||||
--flow_shift 3
|
||||
--seed 1000
|
||||
)
|
||||
|
||||
# DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,757,522'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--generator_update_interval 5
|
||||
--real_score_guidance_scale 3.5
|
||||
--VSA_sparsity 0.9
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/wan_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
@@ -0,0 +1,137 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --nodes=8
|
||||
#SBATCH --ntasks=8
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=dmd_t2v_output/t2v_%j.out
|
||||
#SBATCH --error=dmd_t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate your_env
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR=your_data_dir
|
||||
VALIDATION_DATASET_FILE=your_validation_dataset_file
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_distill_dmd
|
||||
--output_dir "checkpoints/wan_t2v_finetune"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 64
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "3"
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 500
|
||||
--weight_only_checkpointing_steps 500
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--ema_start_step 0
|
||||
--flow_shift 8
|
||||
--seed 1000
|
||||
)
|
||||
|
||||
# DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,757,522'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--generator_update_interval 5
|
||||
--real_score_guidance_scale 3.5
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/training/wan_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "FastVideo/Wan-Syn_77x448x832_600k" --local_dir "FastVideo/Wan-Syn_77x448x832_600k" --repo_type "dataset"
|
||||
@@ -0,0 +1,516 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "In the video, a woman is elegantly showcasing her earrings, bringing attention to their intricate design with a gentle touch of her fingers. She is bathed in ambient purple and pink lighting, which casts a soft glow on her delicate features and enhances the vivid tones of her lipstick and eye makeup. Her hair is styled to frame her face smoothly, emphasizing the contours of her jawline and cheekbones. The background features a blurred neon light, adding an artistic and modern touch to the overall aesthetic.",
|
||||
"video_path": "Fashion/mixkit-face-of-an-elegant-and-captivating-woman-41914_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the video, a lone rider guides a majestic horse across an expansive, open field as the sun sets in the background. The rider, dressed in a classic blue shirt and wide-brimmed hat, sits confidently in the saddle, silhouetted against the warm glow of the evening sky. The horse moves gracefully, its mane and tail flowing with each step, creating a sense of harmony between horse and rider. Surrounding the pair, towering trees form a natural border, their leaves gently rustling in the breeze. The shadows lengthen on the ground, accentuating the serene and timeless feel of the scene. The distant hills and wooden fences frame the horizon, adding depth to the tranquil landscape. A few horses graze peacefully in the background, blending into the pastoral setting. The overall ambiance evokes a sense of calmness and quietude, capturing a perfect moment in the golden light of dusk.",
|
||||
"video_path": "Man/mixkit-a-rancher-riding-a-horse-at-sunset-1143_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a dimly lit, eerie setting, a mysterious pink bottle labeled \"Authentic 100% organic POISON\" sits prominently in the foreground, casting a menacing aura. The bottle is accentuated by green fog, which swirls lightly around it, enhancing its sinister allure. Behind it, a shadowy golden bottle adorned with a spider emblem subtly emerges, adding an extra layer of mystery to the scene. Dim candles provide faint, flickering light, which complements the dark atmosphere, making the setting ideal for an illusion of hidden dangers.",
|
||||
"video_path": "smoke/mixkit-poison-in-halloween-ritual-33879_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The video opens with a tranquil scene in the heart of a dense forest, emphasizing two large, textured tree trunks in the foreground framing the view. Sunlight filters through the canopy above, casting intricate patterns of light and shadow on the trees and the ground. Between the tree trunks, a clear view of a calm, muddy river unfolds, its surface shimmering under the gentle sunlight. The riverbank is decorated with a variety of small bushes and vibrant foliage, subtly transitioning into the deep greens of tall, leafy plants. In the background, the dense forest looms, filled with dark, towering trees, their branches intertwining to form an intricate canopy. The scene is bathed in the soft glow of the sun, creating a serene and picturesque setting. Occasional sunbeams pierce through the foliage, adding a magical aura to the landscape. The vibrant reds and oranges of the smaller plants add contrast, bringing warmth to the earthy tones of the scenery. Overall, this harmonious blend of natural elements creates a peaceful and idyllic forest setting.",
|
||||
"video_path": "forest/mixkit-view-of-a-river-between-two-old-trees-560_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the video, a martial artist dressed in a traditional white uniform with a black belt demonstrates a series of precise movements against a stark black background. The individual gracefully transitions between stances, embodying a sense of focused discipline and control. Each motion is executed with a deliberate pace, showcasing the fluidity of martial arts techniques. The soft lighting creates subtle highlights on the uniform, adding depth to the figure as it moves. The practitioner begins with an open-hand pose, feet firmly grounded, gradually shifting to a powerful forward punch. The fluidity of the sequence displays a mastery of balance and poise. Every trajectory of the limbs is precise and deliberate, capturing the elegance and strength of martial arts. The serene, isolated setting enhances the intensity and concentration of the practitioner. This visual presentation is an elegant interplay of motion and stillness, displaying the art form's discipline and grace.",
|
||||
"video_path": "Man/mixkit-a-young-man-practicing-his-karate-moves-49635_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A tranquil coastal scene unfolds with a drone's aerial view capturing a serene beach landscape. The camera glides over a quiet stretch of sandy shoreline, where gentle waves kiss the shore under a clear blue sky. Nestled amidst lush palm trees are a series of traditional thatched-roof huts, their earthy tones blending harmoniously with the natural surroundings. The sandy beach stretches endlessly, bordered by the rhythmic dance of ocean waves on one side and verdant greenery on the other. A pair of white umbrellas is set up on the sand, suggesting a place to relax and enjoy the sun. In the distance, two small human figures can be seen walking leisurely along the water's edge, leaving faint footprints behind them. The scene exudes a calm and inviting atmosphere, with the soft rustle of palm leaves and the whisper of the ocean breeze almost audible. The overall composition is a captivating blend of nature's tranquility and architectural simplicity. This picturesque setting invites viewers to imagine themselves steps away from this idyllic coastal escape.",
|
||||
"video_path": "beach/mixkit-sunny-beach-in-a-dynamic-shot-from-a-drone-44383_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A lone figure stands on a large, moss-covered rock, surrounded by the soft rush of a nearby stream. The figure is wearing white sneakers and shorts, with a plaid shirt that hangs loosely in the breeze. The lighting creates dramatic shadows, enhancing the textures of the rock and the subtle movement of the water below. In the background, a waterfall cascades into the stream, completing this tranquil and serene nature scene.",
|
||||
"video_path": "forest/mixkit-woman-standing-in-front-of-waterfall-559_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In an industrial setting, a person leans casually against a railing, exuding a sense of confidence and composure. They are wearing a striking outfit, consisting of a vibrant, patterned jacket over a simple white crop top, creating a bold contrast. The atmosphere is infused with warm, ambient lighting that casts soft shadows on the concrete walls and metallic surfaces. Intricate wiring and pipes form an intricate backdrop, enhancing the urban aesthetic. Their relaxed posture and direct, engaging gaze suggest a sense of ease in this industrial environment. This scene encapsulates a blend of modern fashion and gritty, urban architecture, creating a visually compelling narrative.",
|
||||
"video_path": "Fashion/mixkit-portrait-of-a-hipster-woman-walking-down-a-stairs-1297_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A man is energetically stretching in an open-air setting, surrounded by rows of vibrant red seats that suggest an amphitheater or outdoor venue. He wears a sleeveless black shirt layered with a hooded vest, emphasizing his athletic build as he engages in a warm-up routine. Behind him, the striking modern architecture of the building features geometric panels, with large sections of glass and overlapping metallic beams creating a dynamic backdrop. The scene captures the contrast between his focused movements and the static, bold design of the structure, while the surrounding greenery adds a touch of nature to the environment. The overall atmosphere is one of preparation and anticipation, with the man appearing determined and ready for an upcoming event or performance.",
|
||||
"video_path": "Sport/mixkit-man-doing-arm-stretches-595_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A young woman is seated on the floor in front of a plush, beige tufted couch, fully engrossed in sorting through a stack of papers. Her dark hair falls loosely past her shoulders, and she wears a green plaid shirt, contributing to the casual yet focused atmosphere. She gently places the papers onto a small round white table, occasionally lifting individual sheets to examine them more closely. Her expression shifts subtly, reflecting concentration and contemplation as she processes the information on the pages. Two small, round nested tables hold her documents, along with a small plant in a gray pot, adding a touch of greenery to the scene. The background features a dark paneled wall, creating a contrasting backdrop for the light-colored furniture. The setting is tranquil and organized, the couch and tables arranged symmetrically, conveying a sense of harmony. A calculator rests on the smaller table, hinting at a task involving calculations or budgeting.",
|
||||
"video_path": "Woman/mixkit-frustrated-woman-throws-paperwork-on-the-floor-4526_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A heavily rusted metal gate stands firmly locked, with two vertical bars joined by a thick, old chain that loops elegantly around them. The chain's texture is coarse and rugged, its surface reflecting varying shades of orange and brown, indicative of years exposed to the elements. At the heart of the chain, a black iron padlock, slightly worn yet imposing, secures the gate, its curves and edges smooth against the aged links. The gate's metalwork is outlined by a backdrop of soft, blurred greenery, suggesting a serene and isolated location beyond the barrier. Tall trees rise in the distance, their trunks and leaves creating a lush, forest-like setting that contrasts with the gate's severe rust. A pathway leads away from the gate, its surface uneven with patches of moss and weathered stone visible in the soft focus, inviting yet inaccessible. The ambiance is quiet and mysterious, with a sense of abandonment hanging subtly in the air, evoking curiosity about what lies beyond. Shadows play across the gate, cast by branches swaying gently in the breeze, adding to the dynamic interaction of light and texture. This scene, rich in detail and atmosphere, captures the viewer's imagination, evoking both the allure of the forbidden and the beauty of decay.",
|
||||
"video_path": "forest/mixkit-rusty-fence-with-a-chain-of-a-property-in-nature-5294_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a serene and softly lit yoga studio, three individuals engage in a yoga session, each performing an upward-facing stretch. The central figure is a woman with shoulder-length brown hair, dressed in a light cropped top and green leggings, her posture reflecting grace and concentration. To her right, another participant, a woman in a purple outfit, mirrors the pose with equal poise. On her left, a person with a bun focuses intently, supported slightly by yoga blocks beneath their hands. The warm-colored wooden floor contrasts soothingly with the soft pastel mural on the back wall, featuring an abstract design and partial visage of a serene face. Natural light floods the space from a large window on the right, where lush greens peek through, adding an element of tranquility. In the corner of the room, a collection of meditation instruments, including a gong and a Buddha statue, subtly frame the peaceful setting. The mood is calm yet focused, as all three participants are deeply engaged in their practice. The scene combines elements of balance, harmony, and a shared journey towards mindfulness. This depiction captures the essence of a yoga session that blends personal growth with collective experience.",
|
||||
"video_path": "People/mixkit-small-group-of-people-doing-yoga-together-43730_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the deep blue expanse of the ocean, two dolphins glide effortlessly, their sleek bodies reflecting the sunlight filtering through the water. The prominent shadows and caustics create a shimmering effect on their skin, capturing the beauty of their natural habitat. Each dolphin moves with a fluid grace, occasionally interacting with gentle nudges, showcasing their playful and social nature. The scene is vibrant and dynamic, with the clear blue background accentuating the dolphins' movements, making it an ideal subject for AI recreation.",
|
||||
"video_path": "sea/mixkit-dolphins-underwater-4133_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the video, a young woman stands against a vibrant graffiti-covered wall, deeply engrossed in her smartphone. Her expression reflects a mix of focus and subtle satisfaction as she interacts with the screen. She wears a black floral-patterned top, which contrasts with the bright, abstract shapes and bold colors of the mural behind her. As she continues to engage with her phone, a series of like count notifications appear on the screen, indicating a growing online appreciation. The wall behind her features a striking mix of geometric and organic shapes, including swirls of teal, orange, and black, with large humanoid figures in a pop-art style. Her long, light-brown hair frames her face, adding a calm, composed aura amidst the lively backdrop. The video captures a blend of contemporary digital interaction and expressive urban art, creating a dynamic yet harmonious scene.",
|
||||
"video_path": "Girl/mixkit-girl-looking-at-the-likes-in-her-post-4914_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A young mother and her baby sit comfortably on a bed, surrounded by an inviting, cozy atmosphere. The woman, wearing a sleeveless top and jeans, is gently engaging with the baby, who is dressed in an adorable animal-print onesie. The child is seated on the bed with colorful toys scattered around, including a plush toy and a board book. The warm glow from a hanging lamp casts a soft light on them, enhancing the serene environment. Pillows are propped up against the headboard, providing a cushioned backdrop as the mother leans slightly over to interact with the baby. A small bottle is visible beside her, suggesting a nurturing setting. Her hand gestures animatedly as she holds up a soft, white cushion with red and blue accents, likely stimulating the baby\u2019s curiosity. Their shared moment is filled with affection and joy, a perfect snapshot of familial bonding.",
|
||||
"video_path": "Baby/mixkit-loving-mother-and-her-baby-playing-with-soft-toys-49966_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A young girl with long brown hair sits at a round wooden table, engrossed in working on her laptop. The laptop screen is a vivid green, suggesting a green screen effect is in use. To her left, a doll dressed in a yellow and white outfit is casually laid on top of some books, adding a playful and innocent touch to the scene. The setting is cozy, with sheer curtains in the background allowing soft natural light to spill into the room. The girl's posture and focused attention on the laptop suggest she is either playing a game or learning something new. This serene and domestic atmosphere is complemented by the slight blur of a dark couch in the foreground, framing the focused activity of the child.",
|
||||
"video_path": "Girl/mixkit-little-girl-doing-homework-on-a-laptop-4757_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "An expansive view of a calm bay reveals a fleet of sailboats, each anchored in a regimented line stretching toward the horizon. The water is a serene blue, reflecting the soft hues of the early morning sky. A gentle breeze is indicated by the subtle ripples trailing behind the boats, while a single, larger vessel cuts a distinct path, leaving a graceful wake in its journey to the open sea. On one side, a cluster of modern high-rise buildings stands, contrasting against the natural simplicity of the water, suggesting a blend of urban and marine life. The distant shoreline is barely visible, softened by the atmospheric perspective, giving a sense of endless waters meeting the sky. The overall mood is peaceful and orderly, with the boats appearing almost as sentinels guarding the expanse of the tranquil bay.",
|
||||
"video_path": "beach/mixkit-flying-backwards-over-the-sea-near-a-coast-50187_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the video, a person is standing in the center of a dark, featureless space, illuminated by a spotlight that emphasizes their presence. The individual is dressed in a traditional martial arts uniform, known as a gi, which is predominantly white with a black belt tied around the waist, indicating a high level of expertise. The background remains pitch black, creating a stark contrast with the brightly lit figure, ensuring complete focus on them. The person's expression is serious and focused, reflecting a deep sense of discipline and concentration. Their hands move gracefully, transitioning through various martial arts stances, demonstrating practiced skill and fluidity. The uniform's crisp fabric folds and subtly reflects the light, further highlighting each precise movement. Despite the simplicity of the environment, the scene is dynamic, with each motion capturing the essence of martial arts practice. The video effectively conveys a sense of calm strength and mastery, making it ideal for an AI to recreate with attention to posture, lighting, and attire.",
|
||||
"video_path": "Sport/mixkit-karate-fighter-bowing-to-the-front-49706_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a dimly lit room bathed in a mix of neon purple and blue lights, a focused individual is seated in a gaming chair. She wears a white hoodie and large headphones with cat ears that glow softly, creating a striking silhouette. Her hands rest on a keyboard, typing swiftly as she concentrates intently on the screen in front of her. The atmosphere exudes a sense of intensity and immersion, with the soft-colored lighting enhancing the futuristic vibe. Her long hair cascades down her shoulders, adding a touch of elegance to the otherwise tech-centric setting. The overall scene captures the essence of a dedicated gamer deeply engaged in her virtual world.",
|
||||
"video_path": "earth/mixkit-a-young-woman-wearing-headphones-with-rgb-lights-suddenly-gets-51621_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "Inside a dimly-lit bus, five individuals are seated along the rows of worn seats, each subtly illuminated by the colorful lights emanating from overhead. On the left, a woman sits with a relaxed posture, her curly hair accented by a patterned scarf, wearing a plaid outfit paired with bright neon socks. Next to her, a person clad in a denim jacket appears deep in thought, resting their head on a hand. Further back, another figure in a bucket hat and oversized yellow attire gazes across the aisle, evoking a sense of introspection. The atmosphere is enriched by the soft glow of red and green lights, bathing the bus interior in an almost surreal ambiance, creating a compelling tableau of urban life.",
|
||||
"video_path": "Music/mixkit-conceptual-urban-fashion-42581_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "An aerial view captures two tennis players on a court, with one dressed in white on the left and another in red on the right. They are mid-game, each poised for action with rackets in hand, accentuated by their strategic positioning at opposite baselines. The court itself is a stark, deep blue, bordered by the vibrant green of the surrounding area, with a dark central net dividing the space. Long shadows stretch dramatically across the ground, suggesting a late afternoon setting. The subtly textured surface of the court contrasts with the crisp, white lines marking its boundaries and sections. This scene creates a vivid, balanced composition, highlighting both the competitive tension and serene atmosphere of the game.",
|
||||
"video_path": "People/mixkit-two-people-playing-tennis-aerial-view-880_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a vibrant, dreamlike setting, a lone figure moves energetically against a backdrop of deep blue and purple hues, casting emotive shadows that ripple with dynamic motion. The figure, almost obscured by a smeared effect, suggests a rhythmic dance or a passionate performance, arms blurred as they sweep through colorful, streaked lighting. A neon glow accentuates their form, particularly highlighting the face which is abstractly illuminated in bursts of orange and red, suggesting intense emotional expression. The scene is dominated by two primary elements \u2013 the figure\u2019s motion and the dramatic lighting, creating a synergy of human emotion and visual spectacle. Swirling trails of light seem to intertwine with the figure, like a visual symphony of movement and color that floods the space. The lighting changes, casting intricate patterns on the figure and the surrounding space, giving the impression of a kaleidoscope in motion. Despite the blurred and abstract portrayal, there is a sense of focus conveyed through the figure\u2019s intent movements, akin to a conductor orchestrating a visual and auditory performance. The environment resonates with an electric energy, suggesting a seamless fusion of art and technology. As the visual drama unfolds, the scene invites viewers to lose themselves in the abstract dance and the play of vivid luminance.",
|
||||
"video_path": "Music/mixkit-dancer-dancing-with-a-light-bar-in-his-hands-42221_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a brightly lit studio, a photographer wearing a denim jacket focuses intently, capturing shots with a professional camera. Facing him, a model stands gracefully, adjusting her long, flowing hair with delicate movements. The scene is characterized by strong contrasts; the model's soft pink attire and gentle gestures complement the rugged, precise demeanor of the photographer. Positioned against a minimalist backdrop, the pair work seamlessly, with the camera\u2019s lens pointed directly at the model, capturing her elegance. The soft, diffused lighting casts a gentle glow on both subjects, creating an airy and ethereal atmosphere perfect for a high-fashion photo shoot.",
|
||||
"video_path": "Fashion/mixkit-professional-photo-session-with-a-young-female-model-41621_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The video showcases a serene, expansive landscape covered with a variety of trees dotting the hills. The hills gently slope across the frame, with patches of dry grass contrasting against the lush green foliage. Tall trees with dense canopies stand elegantly, casting soft shadows on the ground below. The sunlight bathes the entire scene, highlighting the varied textures of the leaves and terrain. Gaps between the trees reveal a narrow dirt path meandering through the hills, suggesting a sense of quiet solitude. The undulating hills extend into the distance, creating depth and a calming sense of vast space. The verdant hues of the leaves contrast with the earthy tones of the hills, enhancing the visual richness. In the background, a faint outline of distant hills can be seen, blurred softly by the atmospheric perspective. This tranquil setting could be efficiently recreated in a virtual environment by focusing on its layered composition, color palette, and natural textures.",
|
||||
"video_path": "forest/mixkit-aerial-panorama-of-a-sunny-mountain-landscape-40846_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A bustling ski slope comes alive with skiers descending a pristine, snow-covered hill, surrounded by towering, snow-draped evergreens. Several figures stand atop the slope, silhouetted against a clear blue sky, preparing to embark on their ski run. The chair lift on the right continuously drops off eager adventurers, adding to the excitement at the hilltop. Each skier, clad in colorful winter gear, carves distinct paths into the textured snow as they weave their way down. The interplay of sunlight and shadows accentuates the myriad tracks etched into the slope, creating a dynamic visual rhythm. The scene captures a vibrant winter wonderland, full of action and the thrill of a perfect ski day.",
|
||||
"video_path": "Car/mixkit-skiers-on-a-snowy-slope-3327_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The scene unfolds within a dimly lit bus, where three young individuals are seated, each absorbed in their unique world. To the left, a person with tied-back hair rests their head on their hand, dressed casually in a jacket and jeans, projecting a relaxed demeanor. Central to the frame is another individual, sitting upright with intense focus, donning a plaid blazer and oversize hoops, enhancing their confident presence. The muted green and red lighting casts an atmospheric glow, adding depth and intrigue to the setting. On the right, a person in a bucket hat and striped shirt leans back, appearing contemplative as they adjust their hat with a nonchalant gesture. The interplay of light and shadow highlights their expressions, creating an intimate and cinematic ambiance. Together, these figures form a cohesive tableau, capturing a moment of introspection amid a bustling yet serene urban environment.",
|
||||
"video_path": "City/mixkit-three-models-posing-to-the-lens-while-on-board-a-42575_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A determined climber is scaling a massive rock face, showcasing exceptional strength and skill. The person, clad in a teal shirt and dark pants, climbs with precision, their movements measured and deliberate. They are secured by climbing gear, which includes ropes and a harness, emphasizing their commitment to safety. The rugged texture of the sandy-colored rock provides an imposing backdrop, adding drama and scale to the climb. In the distance, other large rock formations and sparse vegetation can be seen under a bright, overcast sky, contributing to the natural and adventurous atmosphere. The scene captures a moment of focus and challenge, highlighting the climber's tenacity and the breathtaking environment.",
|
||||
"video_path": "Sport/mixkit-alpinist-climbing-a-huge-rock-in-a-desert-43306_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A woman stands confidently in front of a large array of solar panels, her navy blue jumpsuit contrasting against the lush green grass beneath her feet. Her expression is calm and focused, eyes facing directly ahead, suggesting a deep connection to the subject matter\u2014renewable energy. The sunlight bathes the scene in warm hues, casting gentle shadows and highlighting the geometric precision of the solar panels' grid-like structure. The background reveals a blend of nature and technology, as the panels are anchored on a grassy slope with foliage on the left side of the frame. This composition captures a harmonious blend of human innovation and environmental consciousness, accentuated by the serene outdoor setting.",
|
||||
"video_path": "Business/mixkit-woman-standing-in-front-of-a-solar-panel-4880_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the video, two people are working at a wooden desk, using an iMac computer. One person, wearing a white knit sweater, is using the apple wireless mouse with their right hand, while their left hand rests on the sleek white keyboard. Their movements are smooth yet intentional, suggesting they are focused on a task on the computer screen. The monitor displays a well-organized array of files and folders, hinting at a task that involves detailed organization or detailed data navigation. The second person, only subtly visible, sits closely by and appears to observe or assist, creating a collaborative atmosphere. Their presence adds a quiet dynamic to the scene, as if they are ready to provide input or guidance. Sticky notes with handwritten notes are attached to the monitor\u2019s stand, adding a touch of personal organization amidst the digital workspace. The focus on the keyboard and mouse emphasizes a streamlined workflow, indicative of a productive work environment. The overall ambiance is calm and focuses on teamwork, technology, and efficient workspace management.",
|
||||
"video_path": "People/mixkit-person-with-glasses-working-on-a-desktop-computer-3248_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A man stands in front of a modern glass facade, taking off a dark hoodie to reveal his gray tank top underneath. His arms are lifted high as he maneuvers the hoodie over his head, showcasing a fluid motion that conveys a sense of calm and routine. The lighting highlights the contours of his muscles, emphasizing a combination of strength and quiet determination. Behind him, the reflective surface of the glass panels provides a subtle backdrop, enhancing the focus on his focused and serene demeanor.",
|
||||
"video_path": "Sport/mixkit-man-puts-on-sleeveless-hoodie-603_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The video displays a captivating dance of fiery orange flames against a stark black background, creating an intense visual contrast. The flames twist and intertwine, forming symmetrical, swirling patterns that expand and contract rhythmically across the frame. Each fiery tendril seems to be alive, moving with an almost hypnotic fluidity that captures the viewer's attention. The illumination from the flames casts subtle shadows, enhancing the depth and texture of the scene. Overall, the dynamic movement and vibrant color palette create an atmosphere of both beauty and power.",
|
||||
"video_path": "fire/mixkit-two-orange-flames-on-black-background-685_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In this scene, a person is seated in a dimly lit room, possibly a recording studio, holding several drumsticks in their hands. The individual's face is partially obscured by sunglasses, adding a touch of mystery to their demeanor. They are wearing a colorful, patterned shirt with a mix of orange and blue tones that stands out against the darker background. The person appears focused and engaged with the drumsticks, their hands prominently displayed. The ambient light casts warm, soft shadows, emphasizing the texture and colors of their shirt and the wooden drumsticks. The room features wooden paneling, which complements the overall cozy, music-centric setting of the scene. The use of perspective centers on the drumsticks, highlighting the importance of rhythm and music in the captured moment.",
|
||||
"video_path": "Music/mixkit-drummer-stretching-before-playing-42783_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A man is casually sitting on a sofa, engrossed in his meal and entertainment. He is holding a TV remote in one hand while reaching for food with the other, indicating a laid-back, comfortable evening. The table before him is filled with takeout containers, revealing a variety of appetizers and dishes, suggestive of a casual dining experience at home. The background is defined by colorful patterned cushions, adding a cozy, homey feel to the scene. Warm, ambient lighting highlights the relaxed atmosphere, casting soft shadows that contribute to the intimate setting. In this moment, he takes a bite of a sandwich, comfortably balancing his attention between food and whatever is playing on the screen.",
|
||||
"video_path": "Man/mixkit-man-watching-tv-and-eating-fast-food-26089_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The scene opens to a breathtaking view of a tranquil ocean horizon at dusk, displaying a vibrant tapestry of oranges, pinks, and purples as the sun sets. In the foreground, tall, swaying palm trees frame the scene, their silhouettes stark against the colorful sky. The ocean itself shimmers with reflections of the sunset, creating a peaceful, almost ethereal atmosphere. A small boat can be seen in the distance, centered on the horizon, adding a sense of scale and solitude to the scene. The waves gently lap the shore, creating faint patterns on the sandy beach, which stretches across the foreground. Above, the sky is dotted with scattered clouds that catch the last light of the day, enhancing the drama and beauty of the scene. The overall mood is serene and contemplative, capturing a perfect moment of nature\u2019s grandeur.",
|
||||
"video_path": "beach/mixkit-sunset-with-sailing-boats-2166_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A man sits hunched on a couch, the weight of emotions clearly visible on his posture. He wears a simple, gray t-shirt, and his head is bowed, resting in his hands, which cover most of his face, obscuring his features. The gentle light filtering through sheer curtains in the background casts a soft glow upon him, emphasizing the contrast between his static form and the hazy brightness behind. His elbows rest upon his knees, suggesting a posture of deep contemplation or distress. The simplicity of the room, with its muted colors, highlights the focus on the man's internal struggle. Delicate detailing on the fabric of his shirt adds texture, enhancing the scene's realism. Subtle changes in the natural light indicate the passage of time, as the man remains unmoving, absorbed in thought. This intimate moment captures a profound vulnerability, making the scene universally relatable and poignant.",
|
||||
"video_path": "Man/mixkit-worried-and-sad-man-with-his-head-down-4701_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A pair of hands, belonging to an unseen figure, carefully unrolls a large sheet of crisp, white paper on a dark wooden table. The lighting is warm, casting a gentle glow that highlights the textures of the paper and the wood grain of the table. As the paper unfurls, the edges reveal the faint beginnings of a colorful map printed on its surface. The arms, clad in a casual gray T-shirt, suggest a relaxed and focused task at hand. Each motion is deliberate, with fingers deftly guiding the paper, ensuring it lays flat without creases. In the background, a hint of a red curtain can be seen, adding a touch of color and depth to the setting. The composition of the scene emphasizes the contrast between the bright paper and the rich tones of the surroundings. This serene and methodical action evokes a sense of exploration and preparation.",
|
||||
"video_path": "Man/mixkit-unrolling-a-world-map-on-a-table-21626_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A young woman sits on a vibrant green seat inside a bus, illuminated by the soft glow of pink and blue lights. Her outfit is a striking mix of colors: a neon pink top paired with a jacket featuring dark sleeves, and jeans that provide a neutral contrast. She wears large, hoop earrings that catch the light as she moves slightly, exuding an air of cool confidence. Her gaze is directed thoughtfully to the side, suggesting contemplation or daydreaming during her commute. The metallic pole beside her adds a geometric element to the composition, reflecting the kaleidoscope of neon hues. The background is a clean, futuristic white, serving as a blank canvas that amplifies the neon atmosphere. Her relaxed posture and the modern bus setting create a scene that captures a blend of urban life and personal introspection.",
|
||||
"video_path": "City/mixkit-fashion-model-posing-on-a-bus-42578_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A silver SUV drives along a winding, snow-covered mountain road, with dense pine trees blanketed in snow lining both sides. The scene is serene, with the vehicle moving smoothly, possibly on a winter journey or vacation. As the SUV disappears around the bend, another, darker SUV follows, creating a sense of motion and perspective on the snow-dusted asphalt. The towering, snow-laden rock formation to the right contrasts with the dark green of the pines, highlighting the peacefulness of the wintry landscape.",
|
||||
"video_path": "Car/mixkit-curve-on-a-snowy-forest-road-3317_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The video showcases a vibrant urban skyline during twilight, with towering buildings reflecting the warm hues of the setting sun. A series of tall, cylindrical structures dominate the foreground, adjacent to a complex of industrial equipment and grids. The scene includes modern high-rise buildings with glass exteriors, capturing the evolving architecture of a bustling cityscape. A prominent structure labeled \"CITY OF AUSTIN POWER PLANT\" stands out, highlighting the industrial theme amidst the urban backdrop. The soft glow of city lights begins to pierce the approaching dusk, creating an inviting yet dynamic atmosphere. Shadows cast by the buildings add depth and contrast, emphasizing their massive scale and intricate designs. The overall composition is balanced between the natural light of the sunset and the artificial illumination of the city, offering a compelling visual narrative.",
|
||||
"video_path": "Car/mixkit-slow-air-travel-in-reverse-over-a-big-city-49841_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the scene, a striking architectural structure dominates the view, bathed in a soft, ambient light. The enormous yellow arches serve as the centerpiece, drawing the eye upwards with their majestic curves and towering presence. The smooth, clean surfaces of the structure reflect the light, highlighting the texture and depth of the architecture. In the foreground, blurred streaks of headlights and taillights suggest the motion of vehicles passing by, adding dynamic energy to the otherwise still scene. The contrast between the fast-moving lights and the static arches creates a balanced composition. To the left, a lone streetlamp and a small tree provide a touch of nature and urban elements against the monumental backdrop. The night sky subtly peeks through the gaps in the structure, hinting at a clear, calm evening. Shadows from the arches create patterns on the ground, adding an intricate detail to the scene. Overall, the combination of light, shadow, and movement makes for a dramatic and visually captivating moment.",
|
||||
"video_path": "Car/mixkit-a-fast-timelapse-of-the-street-with-a-monumental-yellow-50993_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A tranquil marina comes into full view under the golden hues of a setting sun. A collection of gleaming yachts and boats are neatly moored, their reflections shimmering softly on the gentle water. The sun's low position casts elongated shadows over the bustling harbor scene, while rolling hillsides surround the distant cityscape. The skyline is interspersed with modern buildings and clusters of residences, adding layers to the vibrant community. At the center, a broad wooden pier juts confidently into the harbor, extending an invitation for leisurely strolls. To the left, various shops and colorful structures line the waterfront, indicating a vibrant coastal economy. The entire atmosphere exudes a serene yet lively charm, balancing the hustle of maritime activity with the peacefulness of the encroaching dusk. It's a scene of calm anticipation, as if the whole place holds its breath before the night's events unfold.",
|
||||
"video_path": "beach/mixkit-harbor-on-a-tourist-coast-with-many-boats-and-yachts-40077_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The video features a confident individual standing atop a structure against a clear blue sky, exuding a sense of freedom and style. The person is clad in a striking yellow button-up shirt tied at the waist, and beneath it, they wear a simple white top that adds to their relaxed yet stylish appearance. Completing the ensemble are high-waisted white jeans paired with a black belt, adding a touch of contrast. Around their neck is a bold red scarf, providing a splash of color and an air of vintage flair. The person's sunglasses, tinted in yellow, reflect the sunlight and contribute to the overall cool and composed demeanor. Their hair is styled elegantly, pulled back with headphones resting over the ears, suggesting they are immersed in music. One hand casually grazes the headphones, while the other rests gently on the railing, grounding the individual in the moment. The scene is an effortless blend of fashion and tranquility, capturing the spirit of sunny, carefree days.",
|
||||
"video_path": "Music/mixkit-standing-woman-listening-to-music-460_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A ballerina gracefully spins and moves across a pink-hued studio, her poised figure accentuated by a shimmering white tutu and bodice. The background, a continuous wash of soft pink, provides a serene and ethereal atmosphere, emphasizing her fluid movements. Her arms extend with elegance, highlighting the delicacy and precision of her ballet pose, while her focused expression adds intensity to the scene. The subtle details of her costume, combined with the pink monochromatic ambiance, create a dreamlike spectacle, ideal for an AI to envision a oneiric dance setting.",
|
||||
"video_path": "Dance/mixkit-portrait-of-a-ballerina-spinning-with-pink-background-40163_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "The scene unfolds with two human figures in the distance, making their way through a serene meadow, thick with tall golden grass swaying gently in the breeze. The sun hangs low in the sky, casting a soft, diffused glow that illuminates the landscape with a warm, ethereal light. These figures, clad in hiking gear, move deliberately, suggesting they're either embarking on or concluding a journey. Their silhouettes contrast against the lush greenery of the surrounding trees, whose branches reach out, framing the horizon. The play of light and shadow among the trees creates a quilt of textures, with each leaf catching a hint of the sun's dying rays. This tranquil setting evokes a sense of calm and adventure, capturing the quintessential beauty of nature\u2019s landscape.",
|
||||
"video_path": "People/mixkit-landscape-in-nature-while-two-people-are-jogging-44348_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A large cargo ship is docked at an industrial port, its white superstructure contrasting with the deep green and yellow of its deck. The foreground is dominated by the calm, deep blue waters of the harbor, which reflect the vessel\u2019s imposing presence. Surrounding the ship, a series of industrial buildings and storage facilities are visible, hinting at the bustling activity of the port. The deck is intricately detailed, featuring an array of pipes, equipment, and railings, showcasing the ship's functionality and purpose. In the background, a paved area with green patches and a few parked vehicles adds to the busy, industrious atmosphere of the scene.",
|
||||
"video_path": "sea/mixkit-empty-cargo-ship-waiting-at-the-port-4209_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A lone climber ascends a towering rock face, clad in a pink shirt and gray pants, displaying a determined and focused expression. The climber navigates the rugged surface, where the texture of the rock is peppered with natural pockets and crevices that offer handholds and footholds. Sunlight casts soft shadows across the cliff, highlighting the intricate patterns and the climber\u2019s strategic movements. The cliff looms high, with sparse vegetation breaking the monotony of the stone, while distant rocky formations form a dramatic backdrop against the clear blue sky. The climber\u2019s gear, including a harness and chalk bag, underscores the adventure and challenge woven into this majestic, vertical journey.",
|
||||
"video_path": "Sport/mixkit-mountaineer-girl-climbing-a-steep-rocky-mountain-41089_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A person is seen in a close-up shot, skillfully adjusting the tuning pegs of a guitar, showcasing a focused and practiced hand. The image is in black and white, highlighting the contrast between the textures of the instrument and the clothing. The individual's shirt, visible in the background, adds a soft, subtle texture, while the dark tones of the guitar neck create depth in the scene. This composition captures a moment of concentration and finesse, perfect for recreating an intimate musical setting.",
|
||||
"video_path": "Music/mixkit-guitarist-playing-so-inspired-black-and-white-shot-44178_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A musician is playing a large brass instrument with the words \"Brass Band\" clearly visible on its bell. The scene is set against a vibrant yellow backdrop, casting a warm glow on the subject. The musician wears a dark cap and a matching suit, adding a formal touch to his attire. He is deeply focused on his performance, with the instrument's intricate tubing adding complexity to the visual composition. The lighting creates dramatic shadows and highlights, emphasizing the musician's expression and the instrument's metallic sheen. This harmonious blend of color and form captures the essence of a live brass band performance.",
|
||||
"video_path": "Music/mixkit-musician-playing-the-trombone-while-dancing-43752_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the video, a lone musician stands gracefully in front of a grand cathedral, playing an accordion while surrounded by the lively water display of a central fountain. Dressed in a casual ensemble, he wears a light-colored shirt, dark pants, and a flat cap that gives him a vintage charm. His posture is relaxed, yet engaged, as he sways gently in rhythm with the music, casting soft shadows on the cobblestone steps beneath him. The backdrop features the cathedral's towering twin spires, with intricate stonework that casts a rich, historical aura around the scene. Sunlight bathes the entire setting, enhancing the golden hues of the cathedral facade and creating a halo-like effect around the musician. The fountain's water jets splash playfully, catching glimmers of light and adding a dynamic element to the tranquil atmosphere. The scene captures a harmonious blend of architectural majesty and human creativity, framed by the clear, azure sky that extends infinitely above. It's a vivid depiction of solitude and artistry, set against a timeless urban landscape.",
|
||||
"video_path": "Music/mixkit-man-plays-an-accordion-in-front-of-a-fountain-630_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the tranquil video, a person sits in a meditative pose on a gentle hillside, silhouetted against the dawning sky. The person is facing the breathtaking sunrise, with their back slightly turned to the viewer, wearing a simple, light-colored shirt. Their right hand rests on their knee, fingers relaxed in a common meditation mudra, symbolizing calmness and peace. The sky, a stunning blend of soft oranges and deep purples, gradually brightens, casting a warm glow over the lush, green landscape. To the left, the outlines of distant urban buildings can be seen against the horizon, adding a contrast between nature and city life. A river reflecting the sky's colors meanders through the scene, lending a serene, flowing dynamic to the landscape. Trees rise and fall gently across the terrain, their leaves rustling only faintly in the morning breeze. The person remains still and focused, embodying a moment of mindfulness and connection with nature. This visual captures a harmonious balance, evoking a sense of tranquility and introspection.",
|
||||
"video_path": "City/mixkit-girl-meditating-in-yoga-pose-at-sunset-4803_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A serene landscape video captures a breathtaking panoramic view of a vast valley covered in a gentle mist. The undulating hills are lush with dense greenery, their rich foliage creating a vibrant border on the left side of the frame. The mist weaves through the landscape like a soft, ethereal blanket, lending a dream-like quality to the scene. In the distance, several mountain peaks emerge, their dark outlines contrasting against the pale blue sky. A few faint, wispy clouds drift lazily across the horizon, complementing the tranquil atmosphere. The sunlight filters through the haze, casting a warm glow and highlighting different textures of the flora. The overall mood is calm and contemplative, inviting the viewer to pause and appreciate nature's untouched beauty. The composition emphasizes depth and expansiveness, drawing attention to the harmony between earth and sky. This captivating scene embodies tranquility, offering a perfect backdrop for meditation or relaxation.",
|
||||
"video_path": "forest/mixkit-flying-over-a-hill-with-a-view-of-the-surrounding-49743_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In this scene, a bearded individual is intently focused on their smartphone, with the sun setting in the background, casting a warm glow across the cityscape. The person, partially visible, is wearing a dark, buttoned shirt that contrasts with the golden hue of the sunset. Their hands are holding the smartphone delicately but purposefully, reflecting a sense of engagement and focus on the screen. The sunlight creates a striking lens flare effect, enhancing the dramatic atmosphere of the moment as it glimmers off the phone\u2019s surface. The surrounding environment hints at an elevated vantage point, providing a panoramic view of the urban landscape below.",
|
||||
"video_path": "City/mixkit-guy-texting-at-sunset-265_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In an expansive, industrial space defined by towering columns and high ceilings, a solitary figure takes center stage. The person, dressed in dark, fitted clothing, assumes a powerful, dynamic stance with one leg bent forward and both arms outstretched in a horizontal arc. Framing this pose are intense flames that engulf their arms, creating a striking visual contrast against the muted tones of the room. The fire forms a brilliant halo of orange and yellow, casting flickering shadows on the weathered walls and worn, tiled floor. This interplay between light and dark showcases the dancer's poise and agility, as they maintain balance amidst the intense heat. Windows line the background, their panes dimly illuminated by the daylight filtering in, adding depth and perspective to the scene. The entire performance evokes a sense of raw energy and elemental mastery, as the figure continues to manipulate the fire in a seamless, mesmerizing display.",
|
||||
"video_path": "fire/mixkit-expert-juggler-doing-tricks-with-a-stick-with-fire-43663_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A man is playing the violin, focused intently on his music. His fingers gracefully dance along the strings, flawlessly executing each note. He holds the violin close to his chin with a sense of familiarity and expertise. The rich, warm tones of the violin reflect in the soft lighting of the room. He wears a dark shirt, and a subtle necklace rests against his chest, adding a personal touch to his attire. The bow moves smoothly across the strings, producing a melody that seems to fill the space with emotion. His expression is one of concentration and passion, immersing himself fully in the performance. The background is softly blurred, bringing the violin's intricate craftsmanship and his precise movements into sharp focus. This serene and intimate moment captures the essence of his musical artistry.",
|
||||
"video_path": "Music/mixkit-fiddler-playing-a-song-639_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the dimly lit parking garage, two figures engage in an impromptu game of soccer. The first person, wearing a light grey shirt and black pants with three white stripes, skillfully maneuvers the ball with precise footwork. The ground is slick with patches of water, reflecting the vibrant neon lights above. A second figure, clad in dark clothing, stands poised in the background, ready to intercept. The space is defined by stark yellow lines and orange safety bollards, adding structure to the chaotic energy of the scene. The soccer ball glides smoothly across the wet floor, kicking up droplets as it passes. Despite the muted colors of the environment, the players' movements are dynamic and full of life. Their shadowy silhouettes dance with the reflecting light, creating a mesmerizing visual interplay. The atmosphere is charged with focus and camaraderie, encapsulating the essence of a late-night urban soccer experience.",
|
||||
"video_path": "Sport/mixkit-player-making-skillful-play-in-a-street-soccer-game-43504_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A lone climber is seen scaling a towering vertical rock face, demonstrating remarkable strength and focus. Dressed in a light-colored shirt and jeans, the climber grips the stone tightly, navigating the rough textures and crevices with precision. The sheer cliff is massive, exhibiting a range of natural hues from light tan to deep gray, accentuating the climber's figure against the vast rocky backdrop. Surrounding the cliff, scattered greenery and rugged terrain provide a sense of wilderness and isolation. The scene portrays a daring ascension requiring concentration and skill, capturing the essence of human endeavor against nature's formidable beauty.",
|
||||
"video_path": "Sport/mixkit-skilled-mountaineer-climbing-a-gigantic-mountain-41083_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In this serene landscape, a lush meadow stretches across the foreground, dotted with vibrant yellow wildflowers swaying gently in the breeze. A towering tree stands majestically on the right side, its branches reaching wide under the bright blue sky filled with fluffy white clouds. On the left, dense trees form a natural corridor leading to the horizon, suggesting a sense of journey and possibility. The richness of the green grass contrasts beautifully with the golden hue of the distant fields, creating a harmonious palette of nature\u2019s colors. The play of light and shadow adds depth and dimension, evoking a tranquil, inviting atmosphere. It's a scene where nature\u2019s beauty simply commands attention, offering a perfect escape into tranquility.",
|
||||
"video_path": "sky/mixkit-countryside-meadow-4075_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A solitary boat glides across the expansive, tranquil expanse of a serene lake. The vessel leaves a gentle wake behind, creating delicate ripples across the mirror-like surface. The water appears a rich shade of teal, seamlessly blending with the sky at the horizon. Silhouettes of distant trees are faintly visible, creating a picturesque backdrop that enhances the solitary journey of the boat. The sky is a calm gradient, shifting from soft oranges near the shore to the pale blues above. In the distance, a few slender poles emerge from the water, remnants of an old structure or natural formation. The mood of the scene is one of peace and solitude, with the boat journeying steadily through the quiet landscape. There is a sense of endless possibilities as the boat moves toward the unseen beyond the frame. The simplicity and stillness of the scene invite contemplation and reflection, encapsulating a perfect moment of quietude on the water.",
|
||||
"video_path": "mountain/mixkit-motorboat-on-a-large-lake-with-turquoise-blue-waters-4996_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a cozy, dimly lit caf\u00e9, a woman sits alone at a rustic wooden table, fully engrossed in her reading. Her dark, wavy hair frames her face as she leans forward over an open book, suggesting deep focus and contemplation. The caf\u00e9\u2019s ambiance is warm, with hanging pendant lights casting a soft glow over the wooden shelves lined with jars and coffee paraphernalia in the background. A small cup of coffee rests just within her reach, alongside a glass dome encasing a solitary pastry, adding a touch of tranquility to the scene. Her casual attire, a denim jacket over a simple shirt, complements the laid-back, comfortable setting of the caf\u00e9. The contrast between her concentrated expression and the bustling, yet subdued caf\u00e9 atmosphere creates a harmonious, serene visual. The overall composition captures a quiet moment of introspection amidst the gentle hum of caf\u00e9 life.",
|
||||
"video_path": "Woman/mixkit-woman-drinking-coffee-in-a-cafe-223_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a vast, deserted landscape under the night sky, a solitary figure stands at a small music setup, illuminated by strategically placed lights. The person is engrossed in playing a keyboard, with various electronic equipment surrounding them, casting soft glows of orange and blue hues across the scene. To the left, a large circular light adds a dramatic focal point, highlighting the intense contrast between the darkness and the lit performance area. This setup, with its minimalistic design and strategic lighting, creates a captivating and easily recognizable scene that merges the serene, expansive backdrop with an intimate, focused music performance.",
|
||||
"video_path": "Music/mixkit-talented-dj-playing-in-a-lonely-desert-42414_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In a bustling urban scene, cars zoom past a weathered building, their blurred motion a testament to the city\u2019s lively pace. The building, with its faded yellow and brown facade, boasts graffiti that speaks of both art and decay, framing the scene with an air of urban grit. A solitary figure stands slightly to the side, clad casually in a gray top and mustard trousers, gazing into the street, seemingly detached from the surrounding flurry. The motion of the traffic creates a dynamic contrast against the static backdrop, emphasizing the relentless movement of the city. As the video progresses, a bright yellow taxi appears, slowing down as it approaches the figure, adding a pop of color to the desaturated hues of the environment. The interaction suggests a routine, a possibly daily exchange between the driver and the pedestrian, hinting at the rhythms of city life. Overhead, a soft, overcast sky casts a diffused light, lending the scene a subdued, timeless quality. Small elements, like the vertical pole cutting through the frame and the distant chatter of urban sounds, complete this vivid tableau of urban existence.",
|
||||
"video_path": "Car/mixkit-morning-in-the-street-time-lapse-1648_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "A young woman sits on a curb in a tranquil park, basking in the golden hue of the setting sun. Beside her, a collie dog rests calmly, its fur illuminated by the warm sunlight, creating a serene glow. The woman's hand gently strokes the dog's back, highlighting the bond and affection between them. Tall trees surround the pair, casting elongated shadows on the leaf-laden ground, adding to the peaceful and intimate ambiance of the scene.",
|
||||
"video_path": "Pets/mixkit-a-woman-pets-a-dog-in-a-park-1562_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the video, a grand, majestic elephant stands in an open, sunlit field, its massive form dominating the scene. The elephant's skin is a tapestry of earthy tones, with rough, textured wrinkles that add character to its already imposing presence. Its trunk, a powerful and flexible appendage, moves gently, swaying as the elephant possibly enjoys the warmth of the day. The background is a blur of greenery, suggesting a lively environment filled with trees and shrubs that provide a natural habitat. Light plays on the elephant's skin, highlighting patches of dust and dirt that give it an authentic wilderness look. The scene captures the tranquility and majesty of this gentle giant in its natural surroundings.",
|
||||
"video_path": "Zoo/mixkit-wet-elephant-in-the-savanna-3663_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
},
|
||||
{
|
||||
"caption": "In the video, a fluffy dog with brown patches is intently engaged with a bright red toy shaped like a fire hydrant, which has a yellow and orange rope attached. The dog's body is relaxed as it lies on a plain white background, concentrating on nudging and playfully biting the toy. Its ears perk up slightly with curiosity, and its eyes are fixated on the toy, suggesting a scene of focused playfulness. The neutral tones of the dog's fur contrast starkly against the vivid red of the toy, creating a visually striking moment.",
|
||||
"video_path": "Pets/mixkit-a-cute-border-collie-dog-play-with-a-fire-street-50662_clip_1.mp4",
|
||||
"num_inference_steps": 3,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_frames": 61
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,12 @@
|
||||
# Wan2.2-5B Distill Example
|
||||
These are end-to-end example scripts for distilling Wan2.2 TI2V 5B model DMD+VSA methods.
|
||||
|
||||
### 0. Make sure you have installed VSA
|
||||
|
||||
```bash
|
||||
cd csrc/attn
|
||||
git submodule update --init --recursive
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
### TODO
|
||||
@@ -0,0 +1,110 @@
|
||||
#!/bin/bash
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
export MASTER_PORT=29500
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=1
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_ti2v/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/validation.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_distill_dmd_VSA
|
||||
--output_dir="checkpoints/wan_t2v_finetune"
|
||||
--max_train_steps=4000
|
||||
--train_batch_size=1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps=1
|
||||
--num_latent_t 31
|
||||
--num_height 704
|
||||
--num_width 1280
|
||||
--num_frames 121
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--training_state_checkpointing_steps=500
|
||||
--weight_only_checkpointing_steps=500
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus 1
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 1
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "3"
|
||||
--validation_guidance_scale "1.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate=1e-5
|
||||
--mixed_precision="bf16"
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--ema_start_step 0
|
||||
--flow_shift 8
|
||||
--seed 1000
|
||||
)
|
||||
|
||||
# DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,757,522'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--generator_update_interval 5
|
||||
--real_score_guidance_scale 3.5
|
||||
--VSA_sparsity 0.8
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/wan_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}"
|
||||
@@ -0,0 +1,24 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_ti2v/"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
--seed 42 \
|
||||
--max_height 704 \
|
||||
--max_width 1280 \
|
||||
--num_frames 121 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--train_fps 24 \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "t2v"
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"num_frames": 121
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"num_frames": 121
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"num_frames": 121
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -18,6 +18,16 @@ git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
|
||||
python examples/inference/basic/basic.py
|
||||
```
|
||||
|
||||
For an example on Apple silicon:
|
||||
```
|
||||
python examples/inference/basic/basic_mps.py
|
||||
```
|
||||
|
||||
For an example running DMD+VSA inference:
|
||||
```
|
||||
python examples/inference/basic/basic_dmd.py
|
||||
```
|
||||
|
||||
## Basic Walkthrough
|
||||
|
||||
All you need to generate videos using multi-gpus from state-of-the-art diffusion pipelines is the following few lines!
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
@@ -8,25 +8,35 @@ def main():
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
model_name = "Wan-AI/Wan2.2-TI2V-5B-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
use_cpu_offload=False
|
||||
dit_cpu_offload=True,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
# sampling_param.num_frames = 45
|
||||
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.image_path = "test.jpg"
|
||||
# sampling_param.num_inference_steps = 0
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
i2v_prompt = "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
|
||||
i2v_prompt = "A little girl is packing a suitcase and the contents starts flying out of the suitcase everywhere."
|
||||
prompt = i2v_prompt
|
||||
# prompt = (
|
||||
# "A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
# "wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
# "natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
# )
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
return
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
|
||||
@@ -0,0 +1,60 @@
|
||||
import os
|
||||
import time
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples_dmd2"
|
||||
def main():
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
|
||||
load_start_time = time.perf_counter()
|
||||
model_name = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers"
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
# Adjust these offload parameters if you have < 32GB of VRAM
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
VSA_sparsity=0.8,
|
||||
)
|
||||
load_end_time = time.perf_counter()
|
||||
load_time = load_end_time - load_start_time
|
||||
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
|
||||
prompt = (
|
||||
"A neon-lit alley in futuristic Tokyo during a heavy rainstorm at night. The puddles reflect glowing signs in kanji, advertising ramen, karaoke, and VR arcades. A woman in a translucent raincoat walks briskly with an LED umbrella. Steam rises from a street food cart, and a cat darts across the screen. Raindrops are visible on the camera lens, creating a cinematic bokeh effect."
|
||||
)
|
||||
prompt = "A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds."
|
||||
start_time = time.perf_counter()
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
|
||||
end_time = time.perf_counter()
|
||||
gen_time = end_time - start_time
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
# model!
|
||||
prompt2 = (
|
||||
"A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
start_time = time.perf_counter()
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=False)
|
||||
end_time = time.perf_counter()
|
||||
gen_time2 = end_time - start_time
|
||||
|
||||
|
||||
print(f"Time taken to load model: {load_time} seconds")
|
||||
print(f"Time taken to generate video: {gen_time} seconds")
|
||||
print(f"Time taken to generate video2: {gen_time2} seconds")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,40 @@
|
||||
from fastvideo import VideoGenerator, PipelineConfig
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
def main():
|
||||
config = PipelineConfig.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
config.text_encoder_precisions = ["fp16"]
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
pipeline_config=config,
|
||||
use_fsdp_inference=False, # Disable FSDP for MPS
|
||||
dit_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=True,
|
||||
disable_autocast=False,
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
# Create sampling parameters with reduced number of frames
|
||||
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
sampling_param.num_frames = 25 # Reduce from default 81 to 25 frames bc we have to use the SDPA attn backend for mps
|
||||
sampling_param.height = 256
|
||||
sampling_param.width = 256
|
||||
|
||||
prompt = ("A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones.")
|
||||
|
||||
video = generator.generate_video(prompt, sampling_param=sampling_param)
|
||||
|
||||
prompt2 = ("A majestic lion strides across the golden savanna, its powerful frame "
|
||||
"glistening under the warm afternoon sun. The tall grass ripples gently in "
|
||||
"the breeze, enhancing the lion's commanding presence. The tone is vibrant, "
|
||||
"embodying the raw energy of the wild. Low angle, steady tracking shot, "
|
||||
"cinematic.")
|
||||
|
||||
video2 = generator.generate_video(prompt2, sampling_param=sampling_param)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,62 +0,0 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
|
||||
def main():
|
||||
|
||||
# This is the config class for the model initialization
|
||||
config = PipelineConfig.from_pretrained("FastVideo/FastHunyuan-Diffusers")
|
||||
# can be used to dump the config to a yaml file
|
||||
config.dump_to_yaml("config.yaml")
|
||||
print(config)
|
||||
# {
|
||||
# 'vae_config': {
|
||||
# 'scale_factor': 8,
|
||||
# 'sp': True,
|
||||
# 'tiling': True,
|
||||
# 'precision': 'fp16'
|
||||
# },
|
||||
# 'text_encoder_config': {
|
||||
# 'precision': 'fp16'
|
||||
# },
|
||||
# 'dit_config': {
|
||||
# 'precision': 'fp16'
|
||||
# },
|
||||
# 'inference_args': {
|
||||
# 'guidance_scale': 7.5,
|
||||
# 'num_inference_steps': 5,
|
||||
# 'seed': 1024,
|
||||
# 'guidance_rescale': 0.0,
|
||||
# 'flow_shift': 17,
|
||||
# 'num_inference_steps': 5,
|
||||
# }
|
||||
# }
|
||||
|
||||
config.vae_config.scale_factor = 16
|
||||
|
||||
# FastVideo will automatically used the optimal default arguments for the model
|
||||
# If a local path is provided, FastVideo will make a best effort attempt to
|
||||
# identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/FastHunyuan-Diffusers",
|
||||
num_gpus=4,
|
||||
config=config,
|
||||
# or
|
||||
config_path="config.yaml",
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained(
|
||||
"FastVideo/FastHunyuan-Diffusers")
|
||||
sampling_param.num_inference_steps = 5
|
||||
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = "A beautiful woman in a red dress walking down a street"
|
||||
video = generator.generate_video(prompt,
|
||||
sampling_param=sampling_param,
|
||||
num_inference_steps=6)
|
||||
|
||||
video2 = generator.generate_video(prompt2)
|
||||
prompt2 = "A beautiful woman in a blue dress walking down a street"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,38 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=False,
|
||||
image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-I2V-14B-480P-Diffusers")
|
||||
sampling_param.num_frames = 61
|
||||
sampling_param.num_inference_steps = 40
|
||||
sampling_param.guidance_scale = 5.0
|
||||
sampling_param.height = 448
|
||||
sampling_param.width = 832
|
||||
sampling_param.seed = 1024
|
||||
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
# Generate videos with the same simple API, regardless of GPU count
|
||||
prompt = (
|
||||
"An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
|
||||
)
|
||||
video = generator.generate_video(prompt, sampling_param=sampling_param, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,182 @@
|
||||
# FastVideo Dual Model Setup (T2V + I2V)
|
||||
|
||||
This document describes the dual model functionality that has been added to the FastVideo Gradio app, supporting both Text-to-Video (T2V) and Image-to-Video (I2V) generation modes with specialized models.
|
||||
|
||||
## Overview
|
||||
|
||||
The Gradio app now supports both Text-to-Video (T2V) and Image-to-Video (I2V) generation modes using specialized models:
|
||||
- **T2V Model**: `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` for text-to-video generation
|
||||
- **I2V Model**: `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers` for image-to-video generation
|
||||
|
||||
Users can switch between modes using the tabbed interface and upload images for I2V generation.
|
||||
|
||||
## Model Configuration
|
||||
|
||||
### T2V Model
|
||||
- **Model**: `FastVideo/FastWan2.1-T2V-1.3B-Diffusers`
|
||||
- **Purpose**: Text-to-video generation
|
||||
- **Default Parameters**: Optimized for text prompts
|
||||
|
||||
### I2V Model
|
||||
- **Model**: `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers`
|
||||
- **Purpose**: Image-to-video generation
|
||||
- **Default Parameters**: Optimized for image animation
|
||||
|
||||
## Changes Made
|
||||
|
||||
### Backend Changes (`ray_serve_backend.py`)
|
||||
|
||||
1. **Added dual model support**:
|
||||
- Separate model paths for T2V and I2V
|
||||
- Automatic model selection based on request type
|
||||
- Independent model initialization
|
||||
|
||||
2. **Updated `VideoGenerationRequest`**:
|
||||
- Added `model_type` field ("t2v" or "i2v")
|
||||
- Added `image_path` field for I2V input
|
||||
|
||||
3. **Enhanced `FastVideoAPI` class**:
|
||||
- Dual model initialization (`t2v_generator` and `i2v_generator`)
|
||||
- Separate default parameters for each model
|
||||
- Automatic model selection in `generate_video` method
|
||||
|
||||
### Frontend Changes (`gradio_frontend.py`)
|
||||
|
||||
1. **Tabbed interface**:
|
||||
- "Text-to-Video" tab for T2V generation
|
||||
- "Image-to-Video" tab for I2V generation
|
||||
|
||||
2. **Automatic model selection**:
|
||||
- T2V tab uses T2V model automatically
|
||||
- I2V tab uses I2V model automatically
|
||||
- Model type sent in API requests
|
||||
|
||||
3. **Separate event handlers**:
|
||||
- `handle_t2v_generation` for text-to-video
|
||||
- `handle_i2v_generation` for image-to-video
|
||||
|
||||
## Usage
|
||||
|
||||
### Starting the Application
|
||||
|
||||
1. **Using the combined startup script (recommended)**:
|
||||
```bash
|
||||
python start_ray_serve_app.py
|
||||
```
|
||||
|
||||
2. **Manual startup**:
|
||||
```bash
|
||||
# Start backend
|
||||
python ray_serve_backend.py \
|
||||
--t2v_model_path "FastVideo/FastWan2.1-T2V-1.3B-Diffusers" \
|
||||
--i2v_model_path "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
|
||||
# Start frontend
|
||||
python gradio_frontend.py --backend_url "http://localhost:8000"
|
||||
```
|
||||
|
||||
### Using T2V Mode
|
||||
|
||||
1. Navigate to the "Text-to-Video" tab
|
||||
2. Enter a text prompt describing the video you want to generate
|
||||
3. Adjust advanced parameters if needed
|
||||
4. Click "Run" to generate the video
|
||||
|
||||
### Using I2V Mode
|
||||
|
||||
1. Navigate to the "Image-to-Video" tab
|
||||
2. Upload an image using the image upload component
|
||||
3. Enter a prompt describing how the image should animate
|
||||
4. Adjust advanced parameters if needed
|
||||
5. Click "Run" to generate the video
|
||||
|
||||
### Example Prompts
|
||||
|
||||
**T2V Examples**:
|
||||
- "A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface."
|
||||
- "A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks."
|
||||
|
||||
**I2V Examples**:
|
||||
- "The image comes to life with subtle movement, the scene gently animating while maintaining the original composition and mood."
|
||||
- "The static image transforms into a dynamic scene with natural motion, preserving the original lighting and atmosphere."
|
||||
|
||||
## Testing
|
||||
|
||||
A comprehensive test script is provided to verify both T2V and I2V functionality:
|
||||
|
||||
```bash
|
||||
python test_i2v.py
|
||||
```
|
||||
|
||||
This script:
|
||||
- Tests backend health
|
||||
- Tests T2V functionality with text prompts
|
||||
- Tests I2V functionality with image uploads
|
||||
- Verifies response formats for both modes
|
||||
- Cleans up test files
|
||||
|
||||
## Technical Details
|
||||
|
||||
### Backend API Changes
|
||||
|
||||
The `/generate_video` endpoint now accepts:
|
||||
|
||||
```json
|
||||
{
|
||||
"prompt": "Animation description",
|
||||
"model_type": "t2v", // or "i2v"
|
||||
"image_path": "/path/to/input/image.png", // for I2V
|
||||
// ... other parameters
|
||||
}
|
||||
```
|
||||
|
||||
### Model Selection Logic
|
||||
|
||||
- **T2V Mode**: Uses `FastVideo/FastWan2.1-T2V-1.3B-Diffusers`
|
||||
- **I2V Mode**: Uses `Wan-AI/Wan2.1-I2V-14B-480P-Diffusers`
|
||||
- **Automatic Selection**: Based on presence of `image_path` and `model_type`
|
||||
|
||||
### Memory Management
|
||||
|
||||
- Both models are loaded independently
|
||||
- Automatic cleanup prevents memory leaks
|
||||
- Temporary files are cleaned up after processing
|
||||
|
||||
## Configuration Options
|
||||
|
||||
### Command Line Arguments
|
||||
|
||||
**Backend**:
|
||||
- `--t2v_model_path`: Path to T2V model
|
||||
- `--i2v_model_path`: Path to I2V model
|
||||
- `--output_path`: Output directory
|
||||
- `--host`, `--port`: Server configuration
|
||||
|
||||
**Frontend**:
|
||||
- `--backend_url`: Backend API URL
|
||||
- `--t2v_model_path`, `--i2v_model_path`: Model paths (for reference)
|
||||
- `--host`, `--port`: Server configuration
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
1. **Model Loading Issues**: Ensure both models are accessible
|
||||
2. **Memory Issues**: The backend includes automatic cleanup
|
||||
3. **Image Upload Failures**: Check image format and size
|
||||
4. **Generation Failures**: Check backend logs for detailed errors
|
||||
|
||||
## Performance Considerations
|
||||
|
||||
- **Model Loading**: Both models are loaded at startup
|
||||
- **Memory Usage**: Higher memory requirements due to dual models
|
||||
- **Generation Time**: I2V may take longer due to larger model size
|
||||
- **GPU Requirements**: Ensure sufficient VRAM for both models
|
||||
|
||||
## Future Enhancements
|
||||
|
||||
Potential improvements:
|
||||
- Model switching without restart
|
||||
- Batch processing for both modes
|
||||
- Advanced image preprocessing
|
||||
- Progress indicators
|
||||
- Result caching and history
|
||||
- Model-specific parameter optimization
|
||||
@@ -0,0 +1,222 @@
|
||||
# FastVideo with Ray Serve Backend and Gradio Frontend
|
||||
|
||||
This setup provides a scalable web application for FastVideo inference using Ray Serve as the backend and Gradio as the frontend.
|
||||
|
||||
## Architecture
|
||||
|
||||
- **Backend**: Ray Serve handles video generation requests with GPU acceleration
|
||||
- **Frontend**: Gradio provides a user-friendly web interface
|
||||
- **Communication**: HTTP REST API between frontend and backend
|
||||
|
||||
## Features
|
||||
|
||||
- ✅ Scalable backend with Ray Serve
|
||||
- ✅ GPU-accelerated video generation
|
||||
- ✅ User-friendly Gradio interface
|
||||
- ✅ Health monitoring and error handling
|
||||
- ✅ All original functionality preserved
|
||||
- ✅ Easy deployment and management
|
||||
|
||||
## Installation
|
||||
|
||||
1. Install the additional dependencies:
|
||||
|
||||
```bash
|
||||
pip install -r requirements_ray_serve.txt
|
||||
```
|
||||
|
||||
2. Ensure you have the FastVideo model available (the default is `FastVideo/FastHunyuan-diffusers`)
|
||||
|
||||
## Usage
|
||||
|
||||
### Option 1: Start Both Services Together (Recommended)
|
||||
|
||||
Use the startup script to launch both backend and frontend:
|
||||
|
||||
```bash
|
||||
python start_ray_serve_app.py
|
||||
```
|
||||
|
||||
This will:
|
||||
- Start the Ray Serve backend on port 8000
|
||||
- Start the Gradio frontend on port 7860
|
||||
- Monitor both services and provide unified logging
|
||||
- Handle graceful shutdown with Ctrl+C
|
||||
|
||||
### Option 2: Start Services Separately
|
||||
|
||||
#### Start Backend Only
|
||||
|
||||
```bash
|
||||
python ray_serve_backend.py --model_path FastVideo/FastHunyuan-diffusers --output_path outputs
|
||||
```
|
||||
|
||||
#### Start Frontend Only
|
||||
|
||||
```bash
|
||||
python gradio_frontend.py --backend_url http://localhost:8000 --model_path FastVideo/FastHunyuan-diffusers
|
||||
```
|
||||
|
||||
## Configuration
|
||||
|
||||
### Command Line Arguments
|
||||
|
||||
#### Startup Script (`start_ray_serve_app.py`)
|
||||
|
||||
- `--model_path`: Path to the FastVideo model (default: `FastVideo/FastHunyuan-diffusers`)
|
||||
- `--output_path`: Directory to save generated videos (default: `outputs`)
|
||||
- `--backend_host`: Backend host to bind to (default: `0.0.0.0`)
|
||||
- `--backend_port`: Backend port (default: `8000`)
|
||||
- `--frontend_host`: Frontend host to bind to (default: `0.0.0.0`)
|
||||
- `--frontend_port`: Frontend port (default: `7860`)
|
||||
- `--skip_backend_check`: Skip backend health check
|
||||
|
||||
#### Backend (`ray_serve_backend.py`)
|
||||
|
||||
- `--model_path`: Path to the FastVideo model
|
||||
- `--output_path`: Directory to save generated videos
|
||||
- `--host`: Host to bind to
|
||||
- `--port`: Port to bind to
|
||||
|
||||
#### Frontend (`gradio_frontend.py`)
|
||||
|
||||
- `--backend_url`: URL of the Ray Serve backend
|
||||
- `--model_path`: Path to the model (for default parameters)
|
||||
- `--host`: Host to bind to
|
||||
- `--port`: Port to bind to
|
||||
|
||||
### Environment Variables
|
||||
|
||||
You can also set these environment variables:
|
||||
|
||||
- `FASTVIDEO_MODEL_PATH`: Path to the FastVideo model
|
||||
- `FASTVIDEO_OUTPUT_PATH`: Directory to save generated videos
|
||||
- `RAY_SERVE_HOST`: Backend host
|
||||
- `RAY_SERVE_PORT`: Backend port
|
||||
- `GRADIO_HOST`: Frontend host
|
||||
- `GRADIO_PORT`: Frontend port
|
||||
|
||||
## API Endpoints
|
||||
|
||||
### Backend API (Ray Serve)
|
||||
|
||||
- `GET /health`: Health check endpoint
|
||||
- `POST /generate_video`: Video generation endpoint
|
||||
|
||||
#### Video Generation Request
|
||||
|
||||
```json
|
||||
{
|
||||
"prompt": "A beautiful sunset over the ocean",
|
||||
"negative_prompt": "blurry, low quality",
|
||||
"use_negative_prompt": true,
|
||||
"seed": 42,
|
||||
"guidance_scale": 7.5,
|
||||
"num_frames": 21,
|
||||
"height": 512,
|
||||
"width": 512,
|
||||
"num_inference_steps": 20,
|
||||
"randomize_seed": false
|
||||
}
|
||||
```
|
||||
|
||||
#### Video Generation Response
|
||||
|
||||
```json
|
||||
{
|
||||
"output_path": "/path/to/generated/video.mp4",
|
||||
"seed": 42,
|
||||
"success": true,
|
||||
"error_message": null
|
||||
}
|
||||
```
|
||||
|
||||
## Deployment
|
||||
|
||||
### Local Development
|
||||
|
||||
1. Start the application:
|
||||
```bash
|
||||
python start_ray_serve_app.py
|
||||
```
|
||||
|
||||
2. Access the frontend at: `http://localhost:7860`
|
||||
3. Access the backend API at: `http://localhost:8000`
|
||||
|
||||
### Production Deployment
|
||||
|
||||
For production deployment, consider:
|
||||
|
||||
1. **Load Balancing**: Use a reverse proxy (nginx, traefik) in front of the services
|
||||
2. **Monitoring**: Add monitoring and logging (Prometheus, Grafana)
|
||||
3. **Scaling**: Configure Ray Serve for horizontal scaling
|
||||
4. **Security**: Add authentication and rate limiting
|
||||
5. **Storage**: Use shared storage for video outputs
|
||||
|
||||
### Docker Deployment
|
||||
|
||||
Create a Dockerfile for containerized deployment:
|
||||
|
||||
```dockerfile
|
||||
FROM python:3.9-slim
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# Install dependencies
|
||||
COPY requirements_ray_serve.txt .
|
||||
RUN pip install -r requirements_ray_serve.txt
|
||||
|
||||
# Copy application files
|
||||
COPY . .
|
||||
|
||||
# Expose ports
|
||||
EXPOSE 8000 7860
|
||||
|
||||
# Start the application
|
||||
CMD ["python", "start_ray_serve_app.py"]
|
||||
```
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### Common Issues
|
||||
|
||||
1. **Backend not starting**: Check GPU availability and Ray installation
|
||||
2. **Frontend can't connect**: Verify backend URL and network connectivity
|
||||
3. **Video generation fails**: Check model path and GPU memory
|
||||
4. **Port conflicts**: Change ports using command line arguments
|
||||
|
||||
### Logs
|
||||
|
||||
- Backend logs are prefixed with `[BACKEND]`
|
||||
- Frontend logs are prefixed with `[FRONTEND]`
|
||||
- Use `--skip_backend_check` if you need to debug startup issues
|
||||
|
||||
### Performance Tuning
|
||||
|
||||
- Adjust `num_replicas` in the Ray Serve deployment for scaling
|
||||
- Configure `max_concurrent_queries` based on GPU memory
|
||||
- Use multiple GPUs by modifying `ray_actor_options`
|
||||
|
||||
## Migration from Original Gradio Demo
|
||||
|
||||
The new setup maintains full compatibility with the original functionality:
|
||||
|
||||
1. All parameters and options are preserved
|
||||
2. The same example prompts are included
|
||||
3. The UI layout and behavior are identical
|
||||
4. Video generation quality is the same
|
||||
|
||||
The main differences are:
|
||||
- Backend processing is now handled by Ray Serve
|
||||
- Better error handling and status monitoring
|
||||
- Scalable architecture for production use
|
||||
- Separation of concerns between frontend and backend
|
||||
|
||||
## Contributing
|
||||
|
||||
To extend this setup:
|
||||
|
||||
1. Add new endpoints to `ray_serve_backend.py`
|
||||
2. Update the frontend in `gradio_frontend.py`
|
||||
3. Modify the startup script if needed
|
||||
4. Update this README with new features
|
||||
@@ -0,0 +1,501 @@
|
||||
import argparse
|
||||
import os
|
||||
import time
|
||||
import json
|
||||
import statistics
|
||||
import asyncio
|
||||
import aiohttp
|
||||
from copy import deepcopy
|
||||
from typing import List, Dict, Any
|
||||
import threading
|
||||
|
||||
import torch
|
||||
|
||||
# All the prompts for stress testing
|
||||
STRESS_TEST_PROMPTS = [
|
||||
"A person reading a book with words that float off the pages and form pictures.",
|
||||
"A person diving into a pool of liquid crystal, creating ripples of light.",
|
||||
"A handheld shot chasing after a group of friends laughing and playing on the beach at sunset.",
|
||||
"A mysterious ancient temple hidden in the jungle.",
|
||||
"A high-speed train navigating a steep descent.",
|
||||
"a toy robot wearing blue jeans and a white t shirt taking a pleasant stroll in Antarctica during a winter storm",
|
||||
"A cheetah accelerating to full speed while chasing its prey.",
|
||||
"A serene orchard is in full bloom, with trees heavy with blossoms and bees buzzing around, darting from flower to flower in a display of natural harmony.",
|
||||
"A little child let out a big yawn",
|
||||
"Subtle reflections of a woman on the window of a train moving at hyper-speed in a Japanese city.",
|
||||
"A truck left along the edge of a cliff, revealing the stunning coastal landscape below with waves crashing against the rocks.",
|
||||
"A red bird transforms into a flag",
|
||||
"A zoom-out from a single leaf on a tree to reveal the entire forest, showcasing the vastness and diversity of the woodland.",
|
||||
"A slow-motion video of a liquid droplet bouncing on a water-repellent surface.",
|
||||
"Static camera shot. A dinasour running near some lions and chasing them away.",
|
||||
"an adorable kangaroo wearing purple overalls and cowboy boots taking a pleasant stroll in Mumbai India during a beautiful sunset",
|
||||
"A zoom-in on an artist's brush touching the canvas, highlighting the texture of the paint and the strokes being made.",
|
||||
"an old man wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a colorful festival",
|
||||
"A woman is ascending to the sky from the ground",
|
||||
"View out a window of a giant strange creature walking in rundown city at night, one single street lamp dimly lighting the area.",
|
||||
"An arc shot around a lone tree in a vast, foggy field at dawn, revealing the changing light and shadows.",
|
||||
"A person sculpting a statue out of a waterfall, the water solidifying under their touch.",
|
||||
"The person's forehead creased with concentration as she worked on a challenging puzzle.",
|
||||
"The person's cheeks flushed with pleasure as she savored a delicious meal.",
|
||||
"Hand-drawn simple line art, a young kid looking up into space with a wondrous expression on his face.",
|
||||
"A crab made of different jewlery is walking on the beach. As it walks, it drops different jewelry pieces like diamonds, pearls, etc",
|
||||
"Gold coins are falling out when elevator door opens",
|
||||
"the scene transitions from huge waves into a snowy mountain at sunset",
|
||||
"a giant cathedral is completely filled with cats. there are cats everywhere you look. a man enters the cathedral and bows before the giant cat king sitting on a throne.",
|
||||
"A mother dog gently picks up a piece of meat and carefully places it in her puppy's bowl, her eyes filled with warmth and care as she watches her little one eat.",
|
||||
"A soap bubble floating in the air, displaying iridescent colors that shift and change as it moves through different angles of light.",
|
||||
"A truck left alongside a train moving through the countryside, matching its speed and revealing the changing landscape.",
|
||||
"An astronaut walking between stone buildings.",
|
||||
"A close-up shot of the person's face reveals his fear and desperation as he navigates the ship through the storm.",
|
||||
"A frozen lake slowly cracking and thawing as spring arrives, with sheets of ice breaking apart and drifting across the surface.",
|
||||
"A FPV shot zooming through a tunnel into a vibrant underwater space.",
|
||||
"a toy robot wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a colorful festival",
|
||||
"A person sips on a smoothie, the cool and fruity flavors refreshing her mouth.",
|
||||
"In a vibrant theater, a magician in dazzling attire stands center stage, pulling a comically oversized rubber chicken from an ornate, old-fashioned box. His costume shimmers under the stage lights, adding to the spectacle. The crowd erupts in laughter and applause, their faces filled with joy and amazement. The magician's expression hints at mischievous delight as he holds up the rubber chicken, his performance bringing cheer to the audience.",
|
||||
"A hamster running on a spinning wheel.",
|
||||
"A quaint village nestled in a valley is surrounded by blooming cherry blossoms, with petals drifting through the air as villagers go about their daily activities, adding life to the scene.",
|
||||
"In a tranquil forest clearing, a sparkling waterfall cascades down into a clear pool, surrounded by lush greenery and flowers, with occasional birds fluttering by.",
|
||||
"A woman beamed with pride as she watched her child perform on stage.",
|
||||
"an adorable kangaroo wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a winter storm",
|
||||
"A man is eating salad",
|
||||
"An Asian girl wearing a bright yellow T-shirt and white pants is Hip-Hop dancing",
|
||||
"nighttime footage of a hermit crab using an incandescent lightbulb as its shell",
|
||||
"a toy robot wearing a green dress and a sun hat taking a pleasant stroll in Antarctica during a beautiful sunset",
|
||||
"A goat operating a food truck, serving gourmet grilled cheese sandwiches to a line of animals.",
|
||||
"Macro shot. Man in an antique scuba helmet with dark glass walking out of a flower",
|
||||
"A bustling train station in the heart of a vibrant city.",
|
||||
"Light filtering through a canopy of autumn leaves, casting warm, dappled patterns of yellow, orange, and red onto the ground.",
|
||||
"Chimneys in the setting sun",
|
||||
"A longboarder accelerating downhill, carving through turns.",
|
||||
"A couple runs through a sudden downpour, laughing and splashing in puddles as they try to find shelter.",
|
||||
"A glass of iced coffee condensing water on the outside, with droplets forming and sliding down the glass in slow motion.",
|
||||
"macro shot of a leaf showing tiny trains moving through its veins",
|
||||
"A corgi wearing sunglasses walks on the beach of a tropical island",
|
||||
"Borneo wildlife on the Kinabatangan River",
|
||||
"A beautiful silhouette animation shows a wolf howling at the moon, feeling lonely, until it finds its pack.",
|
||||
"an adorable kangaroo wearing blue jeans and a white t shirt taking a pleasant stroll in Johannesburg South Africa during a colorful festival",
|
||||
"A green monster made of plants walks through an airport.",
|
||||
"A close up view of a glass sphere that has a zen garden within it. There is a small dwarf in the sphere who is raking the zen garden and creating patterns in the sand.",
|
||||
"A person on a scooter colliding with a park bench, the scooter tipping over.",
|
||||
"A tilt-up from a city street, ascending to show the skyline with its mix of modern and historic architecture.",
|
||||
"A chef tossing a pancake into the air and catching it.",
|
||||
"A woman whispering a secret into a friend's ear.",
|
||||
"A vulture circling high in the sky.",
|
||||
"A medieval castle overlooking a bustling renaissance fair.",
|
||||
"a toy robot wearing purple overalls and cowboy boots taking a pleasant stroll in Mumbai India during a beautiful sunset",
|
||||
"A man standing in front of a burning building giving the 'thumbs up' sign.",
|
||||
"The person's cheeks flushed with embarrassment as he told a funny story.",
|
||||
"Llamas and Emus are playing chess",
|
||||
"A woman sipping a steaming cup of tea.",
|
||||
"A tree root bursting through the seat of an ancient, weathered bench, intertwining with the wood.",
|
||||
"Smoke rises from the chimney of a cozy log cabin nestled in the woods, with soft light glowing from the windows, suggesting a warm and inviting atmosphere.",
|
||||
"A close-up of sparkling water being poured into a glass, capturing the detailed flow and bubbles.",
|
||||
"a woman wearing blue jeans and a white t shirt taking a pleasant stroll in Antarctica during a beautiful sunset",
|
||||
"The Glenfinnan Viaduct is a historic railway bridge in Scotland, UK, that crosses over the west highland line between the towns of Mallaig and Fort William. It is a stunning sight as a steam train leaves the bridge, traveling over the arch-covered viaduct. The landscape is dotted with lush greenery and rocky mountains, creating a picturesque backdrop for the train journey. The sky is blue and the sun is shining, making for a beautiful day to explore this majestic spot.",
|
||||
"A piece of elastic fabric being pulled and stretched, then returning to its original size when the tension is released.",
|
||||
"a woman wearing a green dress and a sun hat taking a pleasant stroll in Antarctica during a beautiful sunset",
|
||||
"A video of a water jet cutting through metal, showing the powerful and precise movement of water.",
|
||||
"Car mirrors and sunsets",
|
||||
"Giant Pandas are eating hot noodles in a Chinese restaurant",
|
||||
"A rally car taking a fast turn on a track",
|
||||
"a toy robot wearing purple overalls and cowboy boots taking a pleasant stroll in Mumbai India during a colorful festival",
|
||||
"A crystal-clear icicle slowly dripping as it melts in the warmth of the midday sun, each drop sparkling as it falls.",
|
||||
"A tilt-down from a chandelier in a grand hall, revealing the ornate decor and people mingling below.",
|
||||
"A man is playing the drums under the water",
|
||||
"A person playing an electric guitar made of lightning, with thunderous sound waves.",
|
||||
"A person floating in a bubble, drifting over a bustling cityscape.",
|
||||
"A tilt-down from a starry night sky, revealing a quiet forest clearing bathed in moonlight.",
|
||||
"A pan right through a dense jungle, moving past lush vegetation and exotic wildlife.",
|
||||
"Close-up of a man eating an apple.",
|
||||
"A low-angle shot of a dancer leaping gracefully into the air, making their movement appear even more dynamic and powerful.",
|
||||
"A woman is search her bag trying to find something.",
|
||||
"A bulldozer clears debris from a demolished building, making way for new construction.",
|
||||
"A man sighed in relief as the doctor delivered the good news.",
|
||||
"A tsunami coming through an alley in Bulgaria, dynamic movement.",
|
||||
"Blooming Flowers",
|
||||
"A push-in through a dense crowd at a festival, moving towards a performer on stage who is captivating the audience.",
|
||||
"A truck right through a tranquil garden, moving past blooming flowers, trees, and a small fountain.",
|
||||
"The person's eyes sparkled with excitement as he greeted a friend.",
|
||||
"A person playing chess with a robot on a floating platform above the ocean.",
|
||||
"A gentle breeze rustles the leaves as someone walks down a serene forest path, sunlight filtering through the trees and shifting patterns on the ground as branches sway.",
|
||||
"A rollercoaster ride from a city to a desert and then to an ice world",
|
||||
"A pan left across an ancient library, moving from shelf to shelf, showcasing rows of leather-bound books.",
|
||||
"A mother otter floating on her back in a river, cradling her pup on her stomach to keep it safe and warm in the gentle current.",
|
||||
"an adorable kangaroo wearing purple overalls and cowboy boots taking a pleasant stroll in Johannesburg South Africa during a colorful festival",
|
||||
"a woman wearing a green dress and a sun hat taking a pleasant stroll in Mumbai India during a colorful festival",
|
||||
"A delicate layer of morning frost melting off a flower petal, the tiny droplets glistening like diamonds in the light.",
|
||||
"A panda is cooking for her child, her child is next to her.",
|
||||
"Macro shot of a man wearing an antique diving helmet with dark glass and a jetpack walking on the veins of a leaf. Realistic style",
|
||||
"an old man wearing purple overalls and cowboy boots taking a pleasant stroll in Johannesburg South Africa during a beautiful sunset",
|
||||
"A girl is unfolding a birthday gift.",
|
||||
"A pencil drawing an architectural plan.",
|
||||
"A handheld camera following a dog running through a park, bouncing and tilting as it captures the dog's joyful exploration.",
|
||||
"A pan left across a serene beach at sunrise, moving from the darkened shore to the brightening horizon.",
|
||||
"A group of people are clapping to celebrate",
|
||||
"Vendors set up stalls at a bustling farmer's market, displaying fresh fruits and vegetables, while people stroll through, selecting produce and enjoying the lively atmosphere.",
|
||||
"A police helicopter hovers above a high-speed chase, guiding officers on the ground to apprehend a suspect.",
|
||||
"A paper origami dragon riding a boat in waves. Realistic style.",
|
||||
"A close-up of a droplet of dew forming on a leaf, capturing the detailed surface tension.",
|
||||
"a toy robot wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a beautiful sunset",
|
||||
"A dry rainbow rose is coming back to life.",
|
||||
"A glass falling off a table and shattering on the floor.",
|
||||
"A marathon runner crossing the finish line after a grueling race.",
|
||||
"A zoom-in on a drop of morning dew on a leaf, showing the reflection of the surrounding world within it.",
|
||||
"A child blowing on hot cocoa to cool it down.",
|
||||
"A squad of futsal players showcasing their skills on an indoor court.",
|
||||
"A princess is brushing her long golden hair in the garden.",
|
||||
"A close-up of a pair of eyes, revealing the subtle emotions and reflections within them.",
|
||||
"A tracking shot of a group of cyclists racing through a forest trail, with trees and foliage rushing by.",
|
||||
"A woman yawning widely at the end of a long day.",
|
||||
"an old man wearing a green dress and a sun hat taking a pleasant stroll in Johannesburg South Africa during a colorful festival",
|
||||
"Hidden within a garden, an ancient fountain trickles with water, surrounded by vibrant flowers and lush greenery that seem to whisper secrets of the past.",
|
||||
"A Chinese man sits at a table and eats noodles with chopsticks",
|
||||
"A pink pig running fast toward the camera in an alley in Tokyo.",
|
||||
"Strange creatures move through a mysterious, foggy marsh, their silhouettes barely visible through the dense mist as they navigate the eerie, otherworldly landscape.",
|
||||
"Tour of an art gallery with many beautiful works of art in different styles.",
|
||||
"FPV flying through a colorful coral lined streets of an underwater suburban neighborhood.",
|
||||
"Aerial view of Santorini during the blue hour, showcasing the stunning architecture of white Cycladic buildings with blue domes. The caldera views are breathtaking, and the lighting creates a beautiful, serene atmosphere.",
|
||||
"Camera zoom out. A couple walking along the beach as the sun sets over the ocean.",
|
||||
"an extreme close up shot of a woman's eye, with her iris appearing as earth",
|
||||
"a woman wearing purple overalls and cowboy boots taking a pleasant stroll in Mumbai India during a colorful festival",
|
||||
"an old man wearing a green dress and a sun hat taking a pleasant stroll in Mumbai India during a winter storm",
|
||||
"an adorable kangaroo wearing blue jeans and a white t shirt taking a pleasant stroll in Antarctica during a winter storm",
|
||||
"A martial artist breaking a board with a powerful punch.",
|
||||
"People gather on a peaceful beach at sunset, a bonfire crackling as they sit around, enjoying the warmth and the sight of the sun dipping below the horizon.",
|
||||
"A close-up of a waterfall, showing the detailed movement of water as it crashes down.",
|
||||
"A child is blowing bubbles",
|
||||
"a woman wearing a green dress and a sun hat taking a pleasant stroll in Johannesburg South Africa during a winter storm",
|
||||
"A wide-angle perspective of a serene lake surrounded by mountains, reflecting the sky and creating a sense of infinite space.",
|
||||
"The person's eyebrows arched in skepticism as she listened to a dubious claim.",
|
||||
"an old man wearing blue jeans and a white t shirt taking a pleasant stroll in Mumbai India during a beautiful sunset",
|
||||
"a woman wearing a green dress and a sun hat taking a pleasant stroll in Johannesburg South Africa during a colorful festival",
|
||||
"a woman wearing purple overalls and cowboy boots taking a pleasant stroll in Antarctica during a colorful festival",
|
||||
"a toy robot wearing blue jeans and a white t shirt taking a pleasant stroll in Antarctica during a colorful festival",
|
||||
"A chef flips a pancake and puts cream on it.",
|
||||
"An astronaut runs on the surface of the moon, the low angle shot shows the vast background of the moon, the movement is smooth and appears lightweight",
|
||||
"A man's face lit up with happiness as he received a heartfelt compliment.",
|
||||
"A futuristic spaceport hums with activity as ships of various shapes and sizes take off and land on multiple platforms, their engines glowing with vibrant colors.",
|
||||
"A person knitting a scarf using beams of light instead of yarn.",
|
||||
"A pedestal up from the edge of a canyon, gradually revealing the expansive landscape and river below.",
|
||||
"a woman wearing purple overalls and cowboy boots taking a pleasant stroll in Johannesburg South Africa during a colorful festival",
|
||||
"an old man wearing blue jeans and a white t shirt taking a pleasant stroll in Johannesburg South Africa during a colorful festival",
|
||||
"A person walking up a staircase made of clouds leading to a floating castle.",
|
||||
"Monks meditate in a serene mountaintop temple, sitting in quiet reflection as the wind gently moves through the surrounding trees, creating a sense of peace and tranquility.",
|
||||
"An aerial shot of a bustling city intersection at rush hour, capturing the organized chaos of cars and pedestrians.",
|
||||
"A pair of hands skillfully knitting a colorful scarf, the yarn winding through their fingers with each stitch.",
|
||||
"Close-up, a Chinese child is eating dumplings",
|
||||
"A kite losing wind and falling to the ground.",
|
||||
"Bioluminescent waves gently wash ashore on a deserted beach, illuminating the sand with each cresting wave as a figure walks along the water's edge, leaving glowing footprints.",
|
||||
"A red panda taking a bite of a pizza",
|
||||
"A close-up shot of a young woman driving a car, looking thoughtful, blurred green forest visible through the rainy car window.",
|
||||
"A high-speed video of a splash created by a stone thrown into a pond.",
|
||||
"A metal rod being bent slightly by a force and then springing back to its original straight shape when the force is removed.",
|
||||
"A hedgehog in a knight's armor, riding a toy horse into a medieval castle.",
|
||||
"A bird made of fresh oranges rushes out of the orange",
|
||||
"A low altitude first person perspective camera tracking shot of a soccer player's feet dribbling the ball on the groud in a soccer field, Sports Videography, Motion Tracking camera shot",
|
||||
"A tranquil island retreat features swaying palm trees and hammocks strung between them, inviting guests to relax and enjoy the serene beauty of the surroundings.",
|
||||
"a spooky haunted mansion, with friendly jack o lanterns and ghost characters welcoming trick or treaters to the entrance, tilt shift photography",
|
||||
"A coconut tree made of dollar bills at sunset, with bills falling off like leaves.",
|
||||
"A motocross bike accelerating out of a tight turn on a dirt track.",
|
||||
"A tranquil Zen garden with a gently flowing stream and koi fish.",
|
||||
"A green monster made of leaves walks through the airport, carrying a suitcase.",
|
||||
"A time-lapse of a frost-covered leaf gradually thawing in the morning sunlight, with tiny water droplets forming and trickling down.",
|
||||
"A woman practicing her archery skills at a range.",
|
||||
"A slow-motion video of ink being injected into a tank of water, creating intricate and beautiful patterns.",
|
||||
"a woman wearing blue jeans and a white t shirt taking a pleasant stroll in Johannesburg South Africa during a winter storm",
|
||||
"The person's forehead creased with worry as he listened to bad news.",
|
||||
"An arc shot around a grand piano being played in an empty concert hall, the motion revealing the intricate details of the instrument.",
|
||||
"A person conducting a symphony of animals in a forest clearing.",
|
||||
"A truck right alongside a flowing river, capturing the movement of the water and the surrounding forest.",
|
||||
"A rocket blasting off from the launch pad, accelerating rapidly into the sky.",
|
||||
"Workers move through a picturesque vineyard during the harvest season, carefully picking grapes and placing them into baskets as the sun bathes the vines in a warm glow.",
|
||||
"A person is eating an ice cream.",
|
||||
"An over-the-shoulder perspective of a chef meticulously plating a dish in a bustling kitchen.",
|
||||
"A man looked away in shame when confronted with his wrongdoing.",
|
||||
"A person is savoring a slice of pizza at a pizzeria."
|
||||
]
|
||||
|
||||
class BackendStressTest:
|
||||
def __init__(self, output_path: str,
|
||||
server_url: str = "http://localhost:8000", max_concurrent: int = 50):
|
||||
self.output_path = output_path
|
||||
self.server_url = server_url
|
||||
self.max_concurrent = max_concurrent
|
||||
|
||||
# Results storage
|
||||
self.results = []
|
||||
self.lock = threading.Lock()
|
||||
|
||||
async def check_health(self) -> bool:
|
||||
"""Check if the Ray Serve backend is healthy"""
|
||||
try:
|
||||
async with aiohttp.ClientSession() as session:
|
||||
async with session.get(f"{self.server_url}/health", timeout=aiohttp.ClientTimeout(total=5)) as response:
|
||||
return response.status == 200
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def build_request_params(self, prompt: str, **kwargs) -> Dict[str, Any]:
|
||||
"""Build request parameters for Ray Serve backend"""
|
||||
# Default parameters matching the Ray Serve backend
|
||||
default_params = {
|
||||
'prompt': prompt,
|
||||
'negative_prompt': None,
|
||||
'use_negative_prompt': False,
|
||||
'seed': 42,
|
||||
'guidance_scale': 7.5,
|
||||
'num_frames': 21,
|
||||
'height': 448,
|
||||
'width': 832,
|
||||
'num_inference_steps': 20,
|
||||
'randomize_seed': True,
|
||||
'return_frames': False # Don't return frames for stress testing to reduce overhead
|
||||
}
|
||||
|
||||
# Override with any provided kwargs
|
||||
for key, value in kwargs.items():
|
||||
if key in default_params:
|
||||
default_params[key] = value
|
||||
|
||||
# Randomize seed if requested
|
||||
if default_params.get('randomize_seed', True):
|
||||
default_params['seed'] = torch.randint(0, 1000000, (1,)).item()
|
||||
|
||||
# Handle negative prompt
|
||||
if not default_params.get('use_negative_prompt', False):
|
||||
default_params['negative_prompt'] = None
|
||||
|
||||
# NEW: Remove keys with None values to avoid sending nulls that may break validation
|
||||
clean_params = {k: v for k, v in default_params.items() if v is not None}
|
||||
return clean_params
|
||||
|
||||
async def test_single_request(self, session: aiohttp.ClientSession, prompt: str, request_id: int) -> Dict[str, Any]:
|
||||
"""Test a single request and measure latency"""
|
||||
start_time = time.time()
|
||||
|
||||
try:
|
||||
# Build request parameters
|
||||
request_params = self.build_request_params(prompt)
|
||||
|
||||
# Make request to Ray Serve backend
|
||||
async with session.post(
|
||||
f"{self.server_url}/generate_video",
|
||||
json=request_params,
|
||||
timeout=aiohttp.ClientTimeout(total=900) # 15 minute timeout for video generation
|
||||
) as response:
|
||||
|
||||
end_time = time.time()
|
||||
latency = end_time - start_time
|
||||
|
||||
if response.status == 200:
|
||||
response_data = await response.json()
|
||||
if response_data.get('success', False):
|
||||
result = {
|
||||
'request_id': request_id,
|
||||
'prompt': prompt,
|
||||
'latency': latency,
|
||||
'status': 'success',
|
||||
'response_time': latency, # Use our own timing
|
||||
'timestamp': start_time,
|
||||
'output_path': response_data.get('output_path', ''),
|
||||
'used_seed': response_data.get('seed', request_params['seed'])
|
||||
}
|
||||
else:
|
||||
result = {
|
||||
'request_id': request_id,
|
||||
'prompt': prompt,
|
||||
'latency': latency,
|
||||
'status': 'error',
|
||||
'error': response_data.get('error_message', 'Unknown backend error'),
|
||||
'timestamp': start_time
|
||||
}
|
||||
else:
|
||||
response_text = await response.text()
|
||||
result = {
|
||||
'request_id': request_id,
|
||||
'prompt': prompt,
|
||||
'latency': latency,
|
||||
'status': 'error',
|
||||
'error': f"HTTP {response.status}: {response_text}",
|
||||
'timestamp': start_time
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
end_time = time.time()
|
||||
latency = end_time - start_time
|
||||
result = {
|
||||
'request_id': request_id,
|
||||
'prompt': prompt,
|
||||
'latency': latency,
|
||||
'status': 'error',
|
||||
'error': str(e),
|
||||
'timestamp': start_time
|
||||
}
|
||||
|
||||
# Thread-safe result storage
|
||||
with self.lock:
|
||||
self.results.append(result)
|
||||
|
||||
return result
|
||||
|
||||
async def run_stress_test(self, num_iterations: int = 1, concurrent_requests: int = None):
|
||||
"""Run the stress test with multiple iterations and concurrent requests"""
|
||||
if concurrent_requests is None:
|
||||
concurrent_requests = self.max_concurrent
|
||||
|
||||
# Check backend health before starting
|
||||
print(f"Testing Ray Serve backend at {self.server_url}...")
|
||||
if not await self.check_health():
|
||||
print(f"❌ Backend is not healthy at {self.server_url}")
|
||||
print("Make sure the Ray Serve backend is running with:")
|
||||
print("python ray_serve_backend.py")
|
||||
return
|
||||
print("✅ Backend is healthy and ready for stress testing")
|
||||
|
||||
print(f"\nStarting stress test with {len(STRESS_TEST_PROMPTS)} prompts")
|
||||
print(f"Running {num_iterations} iteration(s) with {concurrent_requests} concurrent requests")
|
||||
print(f"Total requests: {len(STRESS_TEST_PROMPTS) * num_iterations}")
|
||||
print(f"Backend URL: {self.server_url}")
|
||||
print("-" * 80)
|
||||
|
||||
all_prompts = STRESS_TEST_PROMPTS * num_iterations
|
||||
request_id = 0
|
||||
|
||||
# Create semaphore to limit concurrent requests
|
||||
semaphore = asyncio.Semaphore(concurrent_requests)
|
||||
|
||||
async def limited_request(session: aiohttp.ClientSession, prompt: str, req_id: int):
|
||||
async with semaphore:
|
||||
return await self.test_single_request(session, prompt, req_id)
|
||||
|
||||
# Run concurrent requests using asyncio
|
||||
async with aiohttp.ClientSession() as session:
|
||||
# Create all tasks
|
||||
tasks = [
|
||||
limited_request(session, prompt, request_id + i)
|
||||
for i, prompt in enumerate(all_prompts)
|
||||
]
|
||||
|
||||
# Process completed requests as they finish
|
||||
completed = 0
|
||||
for coro in asyncio.as_completed(tasks):
|
||||
try:
|
||||
result = await coro
|
||||
completed += 1
|
||||
prompt = result['prompt']
|
||||
status_icon = "✅" if result['status'] == 'success' else "❌"
|
||||
output_info = f" -> {result.get('output_path', 'N/A')}" if result['status'] == 'success' else ""
|
||||
print(f"{status_icon} [{completed}/{len(all_prompts)}] {result['latency']:.2f}s - {prompt[:50]}...{output_info}")
|
||||
except Exception as e:
|
||||
completed += 1
|
||||
print(f"❌ [{completed}/{len(all_prompts)}] Exception: {e}")
|
||||
|
||||
self.analyze_results()
|
||||
|
||||
def analyze_results(self):
|
||||
"""Analyze and print test results"""
|
||||
print("\n" + "=" * 80)
|
||||
print("STRESS TEST RESULTS")
|
||||
print("=" * 80)
|
||||
|
||||
successful_requests = [r for r in self.results if r['status'] == 'success']
|
||||
failed_requests = [r for r in self.results if r['status'] == 'error']
|
||||
|
||||
print(f"Total Requests: {len(self.results)}")
|
||||
print(f"Successful: {len(successful_requests)}")
|
||||
print(f"Failed: {len(failed_requests)}")
|
||||
print(f"Success Rate: {len(successful_requests)/len(self.results)*100:.1f}%")
|
||||
|
||||
if successful_requests:
|
||||
latencies = [r['latency'] for r in successful_requests]
|
||||
print(f"\nLatency Statistics (seconds):")
|
||||
print(f" Min: {min(latencies):.2f}")
|
||||
print(f" Max: {max(latencies):.2f}")
|
||||
print(f" Mean: {statistics.mean(latencies):.2f}")
|
||||
print(f" Median: {statistics.median(latencies):.2f}")
|
||||
print(f" Std Dev: {statistics.stdev(latencies):.2f}")
|
||||
|
||||
# Percentiles
|
||||
sorted_latencies = sorted(latencies)
|
||||
p50 = sorted_latencies[int(len(sorted_latencies) * 0.5)]
|
||||
p90 = sorted_latencies[int(len(sorted_latencies) * 0.9)]
|
||||
p95 = sorted_latencies[int(len(sorted_latencies) * 0.95)]
|
||||
p99 = sorted_latencies[int(len(sorted_latencies) * 0.99)]
|
||||
|
||||
print(f" P50: {p50:.2f}")
|
||||
print(f" P90: {p90:.2f}")
|
||||
print(f" P95: {p95:.2f}")
|
||||
print(f" P99: {p99:.2f}")
|
||||
|
||||
if failed_requests:
|
||||
print(f"\nFailed Requests ({len(failed_requests)}):")
|
||||
for req in failed_requests[:5]: # Show first 5 failures
|
||||
print(f" - {req['error']}")
|
||||
if len(failed_requests) > 5:
|
||||
print(f" ... and {len(failed_requests) - 5} more")
|
||||
|
||||
# Save detailed results
|
||||
results_file = os.path.join(self.output_path, "stress_test_results.json")
|
||||
os.makedirs(self.output_path, exist_ok=True)
|
||||
|
||||
with open(results_file, 'w') as f:
|
||||
json.dump({
|
||||
'summary': {
|
||||
'total_requests': len(self.results),
|
||||
'successful_requests': len(successful_requests),
|
||||
'failed_requests': len(failed_requests),
|
||||
'success_rate': len(successful_requests)/len(self.results)*100 if self.results else 0
|
||||
},
|
||||
'latency_stats': {
|
||||
'min': min(latencies) if successful_requests else 0,
|
||||
'max': max(latencies) if successful_requests else 0,
|
||||
'mean': statistics.mean(latencies) if successful_requests else 0,
|
||||
'median': statistics.median(latencies) if successful_requests else 0,
|
||||
'std_dev': statistics.stdev(latencies) if len(successful_requests) > 1 else 0
|
||||
},
|
||||
'detailed_results': self.results
|
||||
}, f, indent=2)
|
||||
|
||||
print(f"\nDetailed results saved to: {results_file}")
|
||||
|
||||
async def main():
|
||||
parser = argparse.ArgumentParser(description="FastVideo Ray Serve Backend Stress Test")
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs",
|
||||
help="Path to save test results")
|
||||
parser.add_argument("--server_url",
|
||||
type=str,
|
||||
default="http://localhost:8000",
|
||||
help="Ray Serve backend URL")
|
||||
parser.add_argument("--max_concurrent",
|
||||
type=int,
|
||||
default=50,
|
||||
help="Maximum concurrent requests")
|
||||
parser.add_argument("--iterations",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of iterations through all prompts")
|
||||
parser.add_argument("--concurrent_requests",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Number of concurrent requests (overrides max_concurrent)")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Create stress test instance
|
||||
stress_test = BackendStressTest(
|
||||
output_path=args.output_path,
|
||||
server_url=args.server_url,
|
||||
max_concurrent=args.max_concurrent
|
||||
)
|
||||
|
||||
# Run the stress test
|
||||
await stress_test.run_stress_test(
|
||||
num_iterations=args.iterations,
|
||||
concurrent_requests=args.concurrent_requests
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
@@ -6,7 +6,7 @@ import gradio as gr
|
||||
import torch
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(description="FastVideo Gradio Demo")
|
||||
|
||||
@@ -0,0 +1,537 @@
|
||||
import argparse
|
||||
import os
|
||||
import requests
|
||||
import json
|
||||
import base64
|
||||
import io
|
||||
from typing import Optional
|
||||
|
||||
import gradio as gr
|
||||
import torch
|
||||
import imageio
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
class RayServeClient:
|
||||
def __init__(self, backend_url: str):
|
||||
self.backend_url = backend_url.rstrip('/')
|
||||
self.generate_endpoint = f"{self.backend_url}/generate_video"
|
||||
self.health_endpoint = f"{self.backend_url}/health"
|
||||
# Default request timeout in seconds. Increase if generation may run longer.
|
||||
self.request_timeout_s = int(os.getenv("FASTVIDEO_GENERATION_TIMEOUT", "900")) # 15 minutes default
|
||||
|
||||
def check_health(self) -> bool:
|
||||
"""Check if the backend is healthy"""
|
||||
try:
|
||||
response = requests.get(self.health_endpoint, timeout=5)
|
||||
return response.status_code == 200
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
def generate_video(self, request_data: dict) -> dict:
|
||||
"""Send video generation request to the backend"""
|
||||
try:
|
||||
response = requests.post(
|
||||
self.generate_endpoint,
|
||||
json=request_data,
|
||||
timeout=self.request_timeout_s,
|
||||
)
|
||||
response.raise_for_status()
|
||||
return response.json()
|
||||
except requests.exceptions.RequestException as e:
|
||||
return {
|
||||
"success": False,
|
||||
"error_message": f"Backend request failed: {str(e)}",
|
||||
"output_path": "",
|
||||
"seed": request_data.get("seed", 42)
|
||||
}
|
||||
|
||||
|
||||
def decode_and_save_video_from_frames(frames_b64: list, output_dir: str, prompt: str, fps: int = 24) -> str:
|
||||
"""Decode base64 frames and save them as a video file"""
|
||||
if not frames_b64:
|
||||
return "No frames to save"
|
||||
|
||||
# Create safe filename from prompt
|
||||
safe_prompt = prompt[:50].replace(' ', '_').replace('/', '_').replace('\\', '_')
|
||||
video_filename = f"{safe_prompt}_frames.mp4"
|
||||
video_path = os.path.join(output_dir, video_filename)
|
||||
|
||||
try:
|
||||
# Decode frames from base64
|
||||
decoded_frames = []
|
||||
|
||||
for i, frame_b64 in enumerate(frames_b64):
|
||||
try:
|
||||
# Remove the data URL prefix if present
|
||||
if frame_b64.startswith('data:image/'):
|
||||
frame_b64 = frame_b64.split(',')[1]
|
||||
|
||||
# Decode base64 to bytes
|
||||
frame_bytes = base64.b64decode(frame_b64)
|
||||
|
||||
# Create PIL Image from bytes
|
||||
image = Image.open(io.BytesIO(frame_bytes))
|
||||
|
||||
# Convert PIL Image to numpy array (same format as video_generator.py)
|
||||
frame_array = np.array(image)
|
||||
decoded_frames.append(frame_array)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to decode frame {i}: {e}")
|
||||
continue
|
||||
|
||||
if not decoded_frames:
|
||||
return "Failed to decode any frames", ""
|
||||
|
||||
# Save as video using imageio (same as video_generator.py)
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
imageio.mimsave(video_path, decoded_frames, fps=fps, format="mp4")
|
||||
|
||||
return f"Saved {len(decoded_frames)} frames as video: {video_path}", video_path
|
||||
|
||||
except Exception as e:
|
||||
return f"Failed to save video: {str(e)}", ""
|
||||
|
||||
|
||||
def create_gradio_interface(backend_url: str, default_params: SamplingParam):
|
||||
"""Create the Gradio interface"""
|
||||
|
||||
# Initialize the Ray Serve client
|
||||
client = RayServeClient(backend_url)
|
||||
|
||||
def generate_video(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed=False,
|
||||
input_image=None,
|
||||
):
|
||||
# Check backend health first
|
||||
if not client.check_health():
|
||||
return None, f"Backend is not available. Please check if Ray Serve is running at {backend_url}", ""
|
||||
|
||||
# Handle input image for I2V
|
||||
image_path = None
|
||||
if input_image is not None:
|
||||
try:
|
||||
# Save the uploaded image to a temporary file
|
||||
import tempfile
|
||||
temp_dir = "temp_images"
|
||||
os.makedirs(temp_dir, exist_ok=True)
|
||||
|
||||
# Generate a unique filename with appropriate extension
|
||||
import uuid
|
||||
# Determine the best format to preserve quality
|
||||
if hasattr(input_image, 'format') and input_image.format:
|
||||
# Use original format if available
|
||||
ext = input_image.format.lower()
|
||||
if ext == 'jpeg':
|
||||
ext = 'jpg'
|
||||
else:
|
||||
# Default to PNG for lossless quality
|
||||
ext = 'png'
|
||||
|
||||
image_filename = f"input_image_{uuid.uuid4().hex[:8]}.{ext}"
|
||||
image_path = os.path.abspath(os.path.join(temp_dir, image_filename))
|
||||
|
||||
# Save the image preserving original quality
|
||||
if ext == 'png':
|
||||
# Use PNG for lossless compression
|
||||
input_image.save(image_path, "PNG", optimize=False)
|
||||
elif ext == 'jpg':
|
||||
# Use high quality JPEG with minimal compression
|
||||
input_image.convert("RGB").save(image_path, "JPEG", quality=95, optimize=False)
|
||||
else:
|
||||
# For other formats, save as PNG to preserve quality
|
||||
input_image.save(image_path, "PNG", optimize=False)
|
||||
|
||||
print(f"Saved input image to: {image_path}")
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to save input image: {e}")
|
||||
image_path = None
|
||||
|
||||
# Prepare request data - always request frames for video creation
|
||||
request_data = {
|
||||
"prompt": prompt,
|
||||
"negative_prompt": negative_prompt,
|
||||
"use_negative_prompt": use_negative_prompt,
|
||||
"seed": seed,
|
||||
"guidance_scale": guidance_scale,
|
||||
"num_frames": num_frames,
|
||||
"height": height,
|
||||
"width": width,
|
||||
"num_inference_steps": num_inference_steps,
|
||||
"randomize_seed": randomize_seed,
|
||||
"return_frames": False, # Always request frames
|
||||
"image_path": image_path,
|
||||
"model_type": "i2v" if image_path else "t2v" # Use I2V model if image is provided, T2V otherwise
|
||||
}
|
||||
|
||||
# Send request to backend
|
||||
response = client.generate_video(request_data)
|
||||
|
||||
# Clean up temporary image file after processing
|
||||
if image_path and os.path.exists(image_path):
|
||||
try:
|
||||
os.remove(image_path)
|
||||
print(f"Cleaned up temporary image: {image_path}")
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to clean up temporary image {image_path}: {e}")
|
||||
|
||||
if response.get("success", False):
|
||||
output_path = response.get("output_path", "")
|
||||
used_seed = response.get("seed", seed)
|
||||
frames_b64 = response.get("frames", [])
|
||||
|
||||
print(f"Used seed: {used_seed}")
|
||||
print(f"Output path: {output_path}")
|
||||
|
||||
# Handle frame extraction and video creation
|
||||
frames_status = ""
|
||||
# if frames_b64:
|
||||
# try:
|
||||
# Get the output directory from the video path
|
||||
# output_dir = os.path.dirname(output_path) if output_path else "outputs"
|
||||
# frames_status, video_path = decode_and_save_video_from_frames(frames_b64, output_dir, prompt)
|
||||
# print(f"Frames: {frames_status}")
|
||||
# except Exception as e:
|
||||
# frames_status = f"Failed to save frames video: {str(e)}"
|
||||
# print(f"Frame extraction error: {e}")
|
||||
# else:
|
||||
# frames_status = "No frames returned from backend"
|
||||
|
||||
# Check if the video file exists
|
||||
if os.path.exists(output_path):
|
||||
return output_path, used_seed, frames_status
|
||||
else:
|
||||
return None, f"Video generated but file not found at {output_path} {frames_status}", frames_status
|
||||
else:
|
||||
error_msg = response.get("error_message", "Unknown error occurred")
|
||||
return None, f"Generation failed: {error_msg}", ""
|
||||
|
||||
# Example prompts
|
||||
examples = [
|
||||
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand's movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
|
||||
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
|
||||
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
|
||||
]
|
||||
|
||||
# Example I2V prompts (for when users upload images)
|
||||
i2v_examples = [
|
||||
"The image comes to life with subtle movement, the scene gently animating while maintaining the original composition and mood.",
|
||||
"The static image transforms into a dynamic scene with natural motion, preserving the original lighting and atmosphere.",
|
||||
"The photograph animates with realistic movement, bringing the frozen moment to life while keeping the original artistic style.",
|
||||
]
|
||||
|
||||
# Create Gradio interface
|
||||
with gr.Blocks() as demo:
|
||||
gr.Markdown("# FastVideo Inference Demo (Ray Serve Backend)")
|
||||
gr.Markdown(f"**Backend URL:** {backend_url}")
|
||||
|
||||
# Backend status indicator
|
||||
status_text = gr.Text(
|
||||
label="Backend Status",
|
||||
value="Checking backend status...",
|
||||
interactive=False
|
||||
)
|
||||
|
||||
def update_status():
|
||||
if client.check_health():
|
||||
return "✅ Backend is healthy and ready"
|
||||
else:
|
||||
return "❌ Backend is not available"
|
||||
|
||||
with gr.Tabs():
|
||||
# Text-to-Video Tab
|
||||
with gr.Tab("Text-to-Video"):
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
prompt = gr.Text(
|
||||
label="Prompt",
|
||||
show_label=False,
|
||||
max_lines=1,
|
||||
placeholder="Enter your prompt",
|
||||
container=False,
|
||||
)
|
||||
run_button = gr.Button("Run", scale=0)
|
||||
|
||||
result = gr.Video(label="Result", show_label=False)
|
||||
error_output = gr.Text(label="Error", visible=False)
|
||||
frames_output = gr.Text(label="Frame Video Status", visible=False)
|
||||
download_file = gr.File(visible=False)
|
||||
|
||||
gr.Examples(examples=examples, inputs=prompt)
|
||||
|
||||
# Image-to-Video Tab
|
||||
with gr.Tab("Image-to-Video"):
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
i2v_prompt = gr.Text(
|
||||
label="Prompt",
|
||||
show_label=False,
|
||||
max_lines=1,
|
||||
placeholder="Describe how the image should animate",
|
||||
container=False,
|
||||
)
|
||||
i2v_run_button = gr.Button("Run", scale=0)
|
||||
|
||||
input_image = gr.Image(
|
||||
label="Input Image",
|
||||
type="pil",
|
||||
show_label=True,
|
||||
container=True,
|
||||
)
|
||||
|
||||
i2v_result = gr.Video(label="Result", show_label=False)
|
||||
i2v_error_output = gr.Text(label="Error", visible=False)
|
||||
i2v_frames_output = gr.Text(label="Frame Video Status", visible=False)
|
||||
i2v_download_file = gr.File(visible=False)
|
||||
|
||||
gr.Examples(examples=i2v_examples, inputs=i2v_prompt)
|
||||
|
||||
# Shared advanced options
|
||||
with gr.Accordion("Advanced options", open=False):
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
height = gr.Slider(
|
||||
label="Height",
|
||||
minimum=256,
|
||||
maximum=1024,
|
||||
step=32,
|
||||
value=448,
|
||||
)
|
||||
width = gr.Slider(
|
||||
label="Width",
|
||||
minimum=256,
|
||||
maximum=1024,
|
||||
step=32,
|
||||
value=832
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
num_frames = gr.Slider(
|
||||
label="Number of Frames",
|
||||
minimum=16,
|
||||
maximum=160,
|
||||
step=16,
|
||||
value=61,
|
||||
)
|
||||
guidance_scale = gr.Slider(
|
||||
label="Guidance Scale",
|
||||
minimum=1,
|
||||
maximum=12,
|
||||
value=3.0,
|
||||
)
|
||||
num_inference_steps = gr.Slider(
|
||||
label="Inference Steps",
|
||||
minimum=3,
|
||||
maximum=100,
|
||||
value=3,
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
use_negative_prompt = gr.Checkbox(
|
||||
label="Use negative prompt", value=False)
|
||||
negative_prompt = gr.Text(
|
||||
label="Negative prompt",
|
||||
max_lines=1,
|
||||
placeholder="Enter a negative prompt",
|
||||
visible=False,
|
||||
)
|
||||
|
||||
seed = gr.Slider(
|
||||
label="Seed",
|
||||
minimum=0,
|
||||
maximum=1000000,
|
||||
step=1,
|
||||
value=1024
|
||||
)
|
||||
randomize_seed = gr.Checkbox(label="Randomize seed", value=False)
|
||||
seed_output = gr.Number(label="Used Seed")
|
||||
|
||||
gr.Examples(examples=examples, inputs=prompt)
|
||||
|
||||
# Event handlers
|
||||
use_negative_prompt.change(
|
||||
fn=lambda x: gr.update(visible=x),
|
||||
inputs=use_negative_prompt,
|
||||
outputs=negative_prompt,
|
||||
)
|
||||
|
||||
def handle_t2v_generation(*args):
|
||||
# For T2V, we pass None as input_image
|
||||
args = list(args)
|
||||
args.append(None) # Add None for input_image
|
||||
result_path, seed_or_error, frames_status = generate_video(*args)
|
||||
|
||||
if result_path and os.path.exists(result_path):
|
||||
# Show frame status if available
|
||||
if frames_status:
|
||||
return (
|
||||
result_path,
|
||||
seed_or_error,
|
||||
gr.update(visible=False), # error_output
|
||||
gr.update(visible=True, value=frames_status), # frames_output
|
||||
gr.update(visible=True, value=result_path) # download_file
|
||||
)
|
||||
else:
|
||||
return (
|
||||
result_path,
|
||||
seed_or_error,
|
||||
gr.update(visible=False), # error_output
|
||||
gr.update(visible=False), # frames_output
|
||||
gr.update(visible=True, value=result_path) # download_file
|
||||
)
|
||||
else:
|
||||
return (
|
||||
None,
|
||||
seed_or_error,
|
||||
gr.update(visible=True, value=seed_or_error), # error_output
|
||||
gr.update(visible=False), # frames_output
|
||||
gr.update(visible=False) # download_file
|
||||
)
|
||||
|
||||
def handle_i2v_generation(*args):
|
||||
# For I2V, we need to reorder args to match generate_video signature
|
||||
# args should be: [i2v_prompt, negative_prompt, use_negative_prompt, seed, guidance_scale, num_frames, height, width, num_inference_steps, randomize_seed, input_image]
|
||||
result_path, seed_or_error, frames_status = generate_video(*args)
|
||||
|
||||
if result_path and os.path.exists(result_path):
|
||||
# Show frame status if available
|
||||
if frames_status:
|
||||
return (
|
||||
result_path,
|
||||
seed_or_error,
|
||||
gr.update(visible=False), # i2v_error_output
|
||||
gr.update(visible=True, value=frames_status), # i2v_frames_output
|
||||
gr.update(visible=True, value=result_path) # i2v_download_file
|
||||
)
|
||||
else:
|
||||
return (
|
||||
result_path,
|
||||
seed_or_error,
|
||||
gr.update(visible=False), # i2v_error_output
|
||||
gr.update(visible=False), # i2v_frames_output
|
||||
gr.update(visible=True, value=result_path) # i2v_download_file
|
||||
)
|
||||
else:
|
||||
return (
|
||||
None,
|
||||
seed_or_error,
|
||||
gr.update(visible=True, value=seed_or_error), # i2v_error_output
|
||||
gr.update(visible=False), # i2v_frames_output
|
||||
gr.update(visible=False) # i2v_download_file
|
||||
)
|
||||
|
||||
# T2V event handler
|
||||
run_button.click(
|
||||
fn=handle_t2v_generation,
|
||||
inputs=[
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed,
|
||||
],
|
||||
outputs=[result, seed_output, error_output, frames_output, download_file],
|
||||
concurrency_limit=20,
|
||||
)
|
||||
|
||||
# I2V event handler
|
||||
i2v_run_button.click(
|
||||
fn=handle_i2v_generation,
|
||||
inputs=[
|
||||
i2v_prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed,
|
||||
input_image,
|
||||
],
|
||||
outputs=[i2v_result, seed_output, i2v_error_output, i2v_frames_output, i2v_download_file],
|
||||
concurrency_limit=20,
|
||||
)
|
||||
|
||||
# Update status periodically
|
||||
demo.load(update_status, outputs=status_text)
|
||||
|
||||
return demo
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="FastVideo Gradio Frontend")
|
||||
parser.add_argument("--backend_url",
|
||||
type=str,
|
||||
default="http://localhost:8000",
|
||||
help="URL of the Ray Serve backend")
|
||||
parser.add_argument("--t2v_model_path",
|
||||
type=str,
|
||||
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
help="Path to the T2V model (for default parameters)")
|
||||
parser.add_argument("--i2v_model_path",
|
||||
type=str,
|
||||
default="Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
help="Path to the I2V model (for default parameters)")
|
||||
parser.add_argument("--host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Host to bind to")
|
||||
parser.add_argument("--port",
|
||||
type=int,
|
||||
default=7860,
|
||||
help="Port to bind to")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Load default parameters from the models
|
||||
# try:
|
||||
default_params = SamplingParam.from_pretrained(args.t2v_model_path)
|
||||
# except Exception as e:
|
||||
# print(f"Warning: Could not load default parameters from {args.t2v_model_path}: {e}")
|
||||
# print("Using fallback default parameters...")
|
||||
# # Create fallback default parameters
|
||||
# default_params = SamplingParam()
|
||||
# default_params.height = 448
|
||||
# default_params.width = 832
|
||||
# default_params.num_frames = 21
|
||||
# default_params.guidance_scale = 7.5
|
||||
# default_params.num_inference_steps = 20
|
||||
# default_params.seed = 1024
|
||||
|
||||
# Create and launch the interface
|
||||
demo = create_gradio_interface(args.backend_url, default_params)
|
||||
|
||||
print(f"Starting Gradio frontend at http://{args.host}:{args.port}")
|
||||
print(f"Backend URL: {args.backend_url}")
|
||||
print(f"T2V Model: {args.t2v_model_path}")
|
||||
print(f"I2V Model: {args.i2v_model_path}")
|
||||
|
||||
demo.queue(max_size=20).launch(
|
||||
server_name=args.host,
|
||||
server_port=args.port,
|
||||
allowed_paths=[os.path.abspath("outputs"), os.path.abspath("temp_images")]
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,182 @@
|
||||
import argparse
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
import threading
|
||||
import signal
|
||||
import requests
|
||||
from pathlib import Path
|
||||
|
||||
# Add the project root to the Python path
|
||||
project_root = Path(__file__).parent.parent.parent.parent
|
||||
sys.path.insert(0, str(project_root))
|
||||
|
||||
|
||||
def check_frontend_health(frontend_url: str, max_retries: int = 30) -> bool:
|
||||
"""Check if the frontend is healthy"""
|
||||
for i in range(max_retries):
|
||||
try:
|
||||
response = requests.get(frontend_url, timeout=5)
|
||||
if response.status_code == 200:
|
||||
print(f"✅ Frontend is healthy at {frontend_url}")
|
||||
return True
|
||||
except requests.exceptions.RequestException:
|
||||
pass
|
||||
|
||||
if i < max_retries - 1:
|
||||
print(f"⏳ Waiting for frontend to start... ({i+1}/{max_retries})")
|
||||
time.sleep(2)
|
||||
|
||||
print(f"❌ Frontend failed to start within {max_retries * 2} seconds")
|
||||
return False
|
||||
|
||||
|
||||
def start_frontend_instance(args, instance_id: int, backend_url: str):
|
||||
"""Start a single frontend instance"""
|
||||
frontend_script = Path(__file__).parent / "gradio_frontend.py"
|
||||
frontend_port = args.frontend_base_port + instance_id
|
||||
|
||||
cmd = [
|
||||
sys.executable, str(frontend_script),
|
||||
"--backend_url", backend_url,
|
||||
"--t2v_model_path", args.t2v_model_path,
|
||||
"--i2v_model_path", args.i2v_model_path,
|
||||
"--host", args.frontend_host,
|
||||
"--port", str(frontend_port)
|
||||
]
|
||||
|
||||
print(f"🎨 Starting Frontend {instance_id + 1} on port {frontend_port}...")
|
||||
print(f"Command: {' '.join(cmd)}")
|
||||
|
||||
# Start the frontend process
|
||||
frontend_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1
|
||||
)
|
||||
|
||||
# Monitor frontend output
|
||||
def monitor_frontend():
|
||||
for line in frontend_process.stdout:
|
||||
print(f"[FRONTEND-{instance_id + 1}] {line.rstrip()}")
|
||||
|
||||
monitor_thread = threading.Thread(target=monitor_frontend, daemon=True)
|
||||
monitor_thread.start()
|
||||
|
||||
return frontend_process, frontend_port
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="FastVideo Multi-Frontend Launcher")
|
||||
|
||||
# Model and output settings
|
||||
parser.add_argument("--t2v_model_path",
|
||||
type=str,
|
||||
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
help="Path to the T2V model")
|
||||
parser.add_argument("--i2v_model_path",
|
||||
type=str,
|
||||
default="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
help="Path to the I2V model")
|
||||
|
||||
# Frontend settings
|
||||
parser.add_argument("--frontend_host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Frontend host to bind to")
|
||||
parser.add_argument("--frontend_base_port",
|
||||
type=int,
|
||||
default=7860,
|
||||
help="Base port for frontend instances")
|
||||
parser.add_argument("--num_frontends",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Number of frontend instances to start")
|
||||
|
||||
# Backend settings
|
||||
parser.add_argument("--backend_url",
|
||||
type=str,
|
||||
default="http://localhost:8000",
|
||||
help="Backend URL for frontends to connect to")
|
||||
|
||||
# Other settings
|
||||
parser.add_argument("--skip_health_check",
|
||||
action="store_true",
|
||||
help="Skip frontend health check")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
print("🎬 FastVideo Multi-Frontend Launcher")
|
||||
print("=" * 50)
|
||||
print(f"T2V Model: {args.t2v_model_path}")
|
||||
print(f"I2V Model: {args.i2v_model_path}")
|
||||
print(f"Backend URL: {args.backend_url}")
|
||||
print(f"Number of Frontends: {args.num_frontends}")
|
||||
print(f"Frontend Base Port: {args.frontend_base_port}")
|
||||
print("=" * 50)
|
||||
|
||||
# Start multiple frontend instances
|
||||
frontend_processes = []
|
||||
frontend_urls = []
|
||||
|
||||
for i in range(args.num_frontends):
|
||||
process, port = start_frontend_instance(args, i, args.backend_url)
|
||||
frontend_processes.append(process)
|
||||
frontend_urls.append(f"http://{args.frontend_host}:{port}")
|
||||
|
||||
# Wait for frontends to be ready
|
||||
if not args.skip_health_check:
|
||||
print("\n⏳ Waiting for frontends to start...")
|
||||
for i, url in enumerate(frontend_urls):
|
||||
if not check_frontend_health(url):
|
||||
print(f"❌ Frontend {i + 1} failed to start. Terminating...")
|
||||
for process in frontend_processes:
|
||||
process.terminate()
|
||||
sys.exit(1)
|
||||
|
||||
print("\n🎉 All frontend instances are starting up!")
|
||||
for i, url in enumerate(frontend_urls):
|
||||
print(f"📺 Frontend {i + 1}: {url}")
|
||||
print("\nPress Ctrl+C to stop all frontend instances...")
|
||||
|
||||
# Signal handler for graceful shutdown
|
||||
def signal_handler(signum, frame):
|
||||
print("\n🛑 Shutting down frontend instances...")
|
||||
for process in frontend_processes:
|
||||
process.terminate()
|
||||
|
||||
# Wait for processes to terminate
|
||||
try:
|
||||
for process in frontend_processes:
|
||||
process.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
print("⚠️ Force killing processes...")
|
||||
for process in frontend_processes:
|
||||
process.kill()
|
||||
|
||||
print("✅ Frontend instances stopped")
|
||||
sys.exit(0)
|
||||
|
||||
signal.signal(signal.SIGINT, signal_handler)
|
||||
signal.signal(signal.SIGTERM, signal_handler)
|
||||
|
||||
# Monitor processes
|
||||
try:
|
||||
while True:
|
||||
# Check if processes are still running
|
||||
for i, process in enumerate(frontend_processes):
|
||||
if process.poll() is not None:
|
||||
print(f"❌ Frontend {i + 1} process died unexpectedly")
|
||||
break
|
||||
|
||||
time.sleep(1)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
signal_handler(signal.SIGINT, None)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,181 @@
|
||||
# Nginx configuration for FastVideo load balancing
|
||||
# This configuration implements the architecture:
|
||||
# ngrok -> nginx reverse proxy -> frontend1/frontend2 -> backend1×8/backend2×8
|
||||
|
||||
events {
|
||||
worker_connections 1024;
|
||||
}
|
||||
|
||||
http {
|
||||
# Basic settings
|
||||
sendfile on;
|
||||
tcp_nopush on;
|
||||
tcp_nodelay on;
|
||||
keepalive_timeout 65;
|
||||
types_hash_max_size 2048;
|
||||
client_max_body_size 100M; # Allow large video uploads
|
||||
|
||||
# Logging
|
||||
access_log /mnt/fast-disks/nfs/hao_lab/FastVideo/outputs/nginx_access.log;
|
||||
error_log /mnt/fast-disks/nfs/hao_lab/FastVideo/outputs/nginx_error.log;
|
||||
|
||||
# Gzip compression
|
||||
gzip on;
|
||||
gzip_vary on;
|
||||
gzip_min_length 1024;
|
||||
gzip_proxied any;
|
||||
gzip_comp_level 6;
|
||||
gzip_types
|
||||
text/plain
|
||||
text/css
|
||||
text/xml
|
||||
text/javascript
|
||||
application/json
|
||||
application/javascript
|
||||
application/xml+rss
|
||||
application/atom+xml
|
||||
image/svg+xml;
|
||||
|
||||
# Upstream for frontend load balancing
|
||||
upstream frontend_servers {
|
||||
# Round-robin load balancing between frontends
|
||||
server 127.0.0.1:7860 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:7861 weight=1 max_fails=3 fail_timeout=30s;
|
||||
upstream frontend_servers {
|
||||
# Round-robin load balancing between frontends
|
||||
server 127.0.0.1:7860 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:7861 weight=1 max_fails=3 fail_timeout=30s;
|
||||
|
||||
# Health check
|
||||
keepalive 32;
|
||||
}
|
||||
|
||||
# Upstream for backend1 load balancing
|
||||
upstream backend1_servers {
|
||||
server 127.0.0.1:8000 weight=1 max_fails=3 fail_timeout=30s; (8 replicas)
|
||||
upstream backend1_servers {
|
||||
# Round-robin load balancing for backend1 replicas
|
||||
server 127.0.0.1:8000 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8001 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8002 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8003 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8004 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8005 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8006 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8007 weight=1 max_fails=3 fail_timeout=30s;
|
||||
|
||||
keepalive 32;
|
||||
}
|
||||
|
||||
# Upstream for backend2 load balancing
|
||||
upstream backend2_servers {
|
||||
server 127.0.0.1:8000 weight=1 max_fails=3 fail_timeout=30s; (8 replicas)
|
||||
upstream backend2_servers {
|
||||
# Round-robin load balancing for backend2 replicas
|
||||
server 127.0.0.1:8010 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8011 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8012 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8013 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8014 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8015 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8016 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8017 weight=1 max_fails=3 fail_timeout=30s;
|
||||
|
||||
keepalive 32;
|
||||
}
|
||||
|
||||
# Main server block
|
||||
server {
|
||||
listen 80;
|
||||
server_name localhost;
|
||||
|
||||
# Security headers
|
||||
add_header X-Frame-Options "SAMEORIGIN" always;
|
||||
add_header X-Content-Type-Options "nosniff" always;
|
||||
add_header X-XSS-Protection "1; mode=block" always;
|
||||
add_header Referrer-Policy "no-referrer-when-downgrade" always;
|
||||
|
||||
# Frontend routes (Gradio interfaces)
|
||||
location / {
|
||||
proxy_pass http://frontend_servers;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
|
||||
# WebSocket support for Gradio
|
||||
proxy_http_version 1.1;
|
||||
proxy_set_header Upgrade $http_upgrade;
|
||||
proxy_set_header Connection "upgrade";
|
||||
|
||||
# Timeouts
|
||||
proxy_connect_timeout 60s;
|
||||
proxy_send_timeout 60s;
|
||||
proxy_read_timeout 60s;
|
||||
|
||||
# Buffer settings
|
||||
proxy_buffering on;
|
||||
proxy_buffer_size 128k;
|
||||
proxy_buffers 4 256k;
|
||||
proxy_busy_buffers_size 256k;
|
||||
}
|
||||
|
||||
# Backend API routes for frontend1
|
||||
location /api/frontend1/ {
|
||||
# Strip the /api/frontend1/ prefix
|
||||
rewrite ^/api/frontend1/(.*) /$1 break;
|
||||
proxy_pass http://backend1_servers;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
|
||||
# Timeouts for video generation
|
||||
proxy_connect_timeout 300s;
|
||||
proxy_send_timeout 300s;
|
||||
proxy_read_timeout 300s;
|
||||
|
||||
# Buffer settings for large responses
|
||||
proxy_buffering on;
|
||||
proxy_buffer_size 128k;
|
||||
proxy_buffers 4 256k;
|
||||
proxy_busy_buffers_size 256k;
|
||||
}
|
||||
|
||||
# Backend API routes for frontend2
|
||||
location /api/frontend2/ {
|
||||
# Strip the /api/frontend2/ prefix
|
||||
rewrite ^/api/frontend2/(.*) /$1 break;
|
||||
proxy_pass http://backend2_servers;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
|
||||
# Timeouts for video generation
|
||||
proxy_connect_timeout 300s;
|
||||
proxy_send_timeout 300s;
|
||||
proxy_read_timeout 300s;
|
||||
|
||||
# Buffer settings for large responses
|
||||
proxy_buffering on;
|
||||
proxy_buffer_size 128k;
|
||||
proxy_buffers 4 256k;
|
||||
proxy_busy_buffers_size 256k;
|
||||
}
|
||||
|
||||
# Health check endpoint
|
||||
location /health {
|
||||
access_log off;
|
||||
return 200 "healthy\n";
|
||||
add_header Content-Type text/plain;
|
||||
}
|
||||
|
||||
# Static files (if needed)
|
||||
location /static/ {
|
||||
alias /var/www/static/;
|
||||
expires 1y;
|
||||
add_header Cache-Control "public, immutable";
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,173 @@
|
||||
# Nginx configuration for FastVideo load balancing
|
||||
# This configuration implements the architecture:
|
||||
# ngrok -> nginx reverse proxy -> frontend1/frontend2 -> backend1×8/backend2×8
|
||||
|
||||
events {
|
||||
worker_connections 1024;
|
||||
}
|
||||
|
||||
http {
|
||||
# Basic settings
|
||||
sendfile on;
|
||||
tcp_nopush on;
|
||||
tcp_nodelay on;
|
||||
keepalive_timeout 65;
|
||||
types_hash_max_size 2048;
|
||||
client_max_body_size 100M; # Allow large video uploads
|
||||
|
||||
# Logging
|
||||
access_log /var/log/nginx/access.log;
|
||||
error_log /var/log/nginx/error.log;
|
||||
|
||||
# Gzip compression
|
||||
gzip on;
|
||||
gzip_vary on;
|
||||
gzip_min_length 1024;
|
||||
gzip_proxied any;
|
||||
gzip_comp_level 6;
|
||||
gzip_types
|
||||
text/plain
|
||||
text/css
|
||||
text/xml
|
||||
text/javascript
|
||||
application/json
|
||||
application/javascript
|
||||
application/xml+rss
|
||||
application/atom+xml
|
||||
image/svg+xml;
|
||||
|
||||
# Upstream for frontend load balancing
|
||||
upstream frontend_servers {
|
||||
# Round-robin load balancing between frontends
|
||||
server 127.0.0.1:7860 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:7861 weight=1 max_fails=3 fail_timeout=30s;
|
||||
|
||||
# Health check
|
||||
keepalive 32;
|
||||
}
|
||||
|
||||
# Upstream for backend1 load balancing (8 replicas)
|
||||
upstream backend1_servers {
|
||||
# Round-robin load balancing for backend1 replicas
|
||||
server 127.0.0.1:8000 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8001 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8002 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8003 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8004 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8005 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8006 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8007 weight=1 max_fails=3 fail_timeout=30s;
|
||||
|
||||
keepalive 32;
|
||||
}
|
||||
|
||||
# Upstream for backend2 load balancing (8 replicas)
|
||||
upstream backend2_servers {
|
||||
# Round-robin load balancing for backend2 replicas
|
||||
server 127.0.0.1:8010 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8011 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8012 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8013 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8014 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8015 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8016 weight=1 max_fails=3 fail_timeout=30s;
|
||||
server 127.0.0.1:8017 weight=1 max_fails=3 fail_timeout=30s;
|
||||
|
||||
keepalive 32;
|
||||
}
|
||||
|
||||
# Main server block
|
||||
server {
|
||||
listen 80;
|
||||
server_name localhost;
|
||||
|
||||
# Security headers
|
||||
add_header X-Frame-Options "SAMEORIGIN" always;
|
||||
add_header X-Content-Type-Options "nosniff" always;
|
||||
add_header X-XSS-Protection "1; mode=block" always;
|
||||
add_header Referrer-Policy "no-referrer-when-downgrade" always;
|
||||
|
||||
# Frontend routes (Gradio interfaces)
|
||||
location / {
|
||||
proxy_pass http://frontend_servers;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
|
||||
# WebSocket support for Gradio
|
||||
proxy_http_version 1.1;
|
||||
proxy_set_header Upgrade $http_upgrade;
|
||||
proxy_set_header Connection "upgrade";
|
||||
|
||||
# Timeouts
|
||||
proxy_connect_timeout 60s;
|
||||
proxy_send_timeout 60s;
|
||||
proxy_read_timeout 60s;
|
||||
|
||||
# Buffer settings
|
||||
proxy_buffering on;
|
||||
proxy_buffer_size 128k;
|
||||
proxy_buffers 4 256k;
|
||||
proxy_busy_buffers_size 256k;
|
||||
}
|
||||
|
||||
# Backend API routes for frontend1
|
||||
location /api/frontend1/ {
|
||||
# Strip the /api/frontend1/ prefix
|
||||
rewrite ^/api/frontend1/(.*) /$1 break;
|
||||
proxy_pass http://backend1_servers;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
|
||||
# Timeouts for video generation
|
||||
proxy_connect_timeout 300s;
|
||||
proxy_send_timeout 300s;
|
||||
proxy_read_timeout 300s;
|
||||
|
||||
# Buffer settings for large responses
|
||||
proxy_buffering on;
|
||||
proxy_buffer_size 128k;
|
||||
proxy_buffers 4 256k;
|
||||
proxy_busy_buffers_size 256k;
|
||||
}
|
||||
|
||||
# Backend API routes for frontend2
|
||||
location /api/frontend2/ {
|
||||
# Strip the /api/frontend2/ prefix
|
||||
rewrite ^/api/frontend2/(.*) /$1 break;
|
||||
proxy_pass http://backend2_servers;
|
||||
proxy_set_header Host $host;
|
||||
proxy_set_header X-Real-IP $remote_addr;
|
||||
proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;
|
||||
proxy_set_header X-Forwarded-Proto $scheme;
|
||||
|
||||
# Timeouts for video generation
|
||||
proxy_connect_timeout 300s;
|
||||
proxy_send_timeout 300s;
|
||||
proxy_read_timeout 300s;
|
||||
|
||||
# Buffer settings for large responses
|
||||
proxy_buffering on;
|
||||
proxy_buffer_size 128k;
|
||||
proxy_buffers 4 256k;
|
||||
proxy_busy_buffers_size 256k;
|
||||
}
|
||||
|
||||
# Health check endpoint
|
||||
location /health {
|
||||
access_log off;
|
||||
return 200 "healthy\n";
|
||||
add_header Content-Type text/plain;
|
||||
}
|
||||
|
||||
# Static files (if needed)
|
||||
location /static/ {
|
||||
alias /var/www/static/;
|
||||
expires 1y;
|
||||
add_header Cache-Control "public, immutable";
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,337 @@
|
||||
import time
|
||||
import os
|
||||
import torch
|
||||
import base64
|
||||
import io
|
||||
from copy import deepcopy
|
||||
from typing import Dict, Any, Optional, List
|
||||
|
||||
import ray
|
||||
from ray import serve
|
||||
from fastapi import FastAPI, Request
|
||||
from pydantic import BaseModel
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
from slowapi import Limiter, _rate_limit_exceeded_handler
|
||||
from slowapi.util import get_remote_address
|
||||
from slowapi.errors import RateLimitExceeded
|
||||
|
||||
|
||||
class VideoGenerationRequest(BaseModel):
|
||||
prompt: str
|
||||
negative_prompt: Optional[str] = None
|
||||
use_negative_prompt: bool = False
|
||||
seed: int = 42
|
||||
guidance_scale: float = 7.5
|
||||
num_frames: int = 21
|
||||
height: int = 448
|
||||
width: int = 832
|
||||
num_inference_steps: int = 20
|
||||
randomize_seed: bool = False
|
||||
return_frames: bool = False # Whether to return base64 encoded frames
|
||||
image_path: Optional[str] = None # Path to input image for I2V
|
||||
model_type: str = "t2v" # "t2v" or "i2v" to specify which model to use
|
||||
|
||||
|
||||
class VideoGenerationResponse(BaseModel):
|
||||
output_path: str
|
||||
seed: int
|
||||
success: bool
|
||||
error_message: Optional[str] = None
|
||||
frames: Optional[List[str]] = None # Base64 encoded frames
|
||||
|
||||
|
||||
def encode_frames_to_base64(frames: List[np.ndarray]) -> List[str]:
|
||||
"""Convert numpy frames (0-255) to base64-encoded PNG images"""
|
||||
if not frames:
|
||||
return []
|
||||
|
||||
encoded_frames = []
|
||||
|
||||
for i, frame in enumerate(frames):
|
||||
try:
|
||||
# Ensure frame is numpy array
|
||||
if not isinstance(frame, np.ndarray):
|
||||
print(f"Warning: Frame {i} is not a numpy array, skipping")
|
||||
continue
|
||||
|
||||
# Ensure frame is uint8
|
||||
if frame.dtype != np.uint8:
|
||||
# Clip values to 0-255 range and convert to uint8
|
||||
frame = np.clip(frame, 0, 255).astype(np.uint8)
|
||||
|
||||
# Convert numpy array to PIL Image
|
||||
if len(frame.shape) == 3 and frame.shape[2] == 3:
|
||||
# RGB image
|
||||
pil_image = Image.fromarray(frame, mode='RGB')
|
||||
elif len(frame.shape) == 3 and frame.shape[2] == 4:
|
||||
# RGBA image
|
||||
pil_image = Image.fromarray(frame, mode='RGBA')
|
||||
elif len(frame.shape) == 2:
|
||||
# Grayscale image
|
||||
pil_image = Image.fromarray(frame, mode='L')
|
||||
else:
|
||||
print(f"Warning: Frame {i} has unsupported shape {frame.shape}, skipping")
|
||||
continue
|
||||
|
||||
# Save to bytes buffer as PNG
|
||||
buffer = io.BytesIO()
|
||||
pil_image.save(buffer, format='PNG')
|
||||
buffer.seek(0)
|
||||
|
||||
# Encode to base64
|
||||
img_base64 = base64.b64encode(buffer.getvalue()).decode('utf-8')
|
||||
encoded_frames.append(f"data:image/png;base64,{img_base64}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to encode frame {i}: {e}")
|
||||
continue
|
||||
|
||||
return encoded_frames
|
||||
|
||||
|
||||
# Create FastAPI app with rate limiting
|
||||
app = FastAPI()
|
||||
|
||||
# Initialize rate limiter
|
||||
limiter = Limiter(key_func=get_remote_address)
|
||||
app.state.limiter = limiter
|
||||
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
|
||||
|
||||
|
||||
@serve.deployment(
|
||||
num_replicas=8,
|
||||
# ray_actor_options={"num_cpus": 10, "num_gpus": 1, "runtime_env": {"conda": "fv", "working_dir": "/mnt/fast-disks/nfs/hao_lab/FastVideo"}},
|
||||
ray_actor_options={"num_cpus": 10, "num_gpus": 1, "runtime_env": {"conda": "fv"}},
|
||||
)
|
||||
@serve.ingress(app)
|
||||
class FastVideoAPI:
|
||||
def __init__(self, t2v_model_path: str, i2v_model_path: str, output_path: str):
|
||||
self.t2v_model_path = t2v_model_path
|
||||
self.i2v_model_path = i2v_model_path
|
||||
self.output_path = output_path
|
||||
|
||||
# Initialize the video generators
|
||||
self.t2v_generator = None # Initialize to None
|
||||
self.i2v_generator = None # Initialize to None
|
||||
self.t2v_default_params = None # Initialize to None
|
||||
self.i2v_default_params = None # Initialize to None
|
||||
|
||||
# Ensure output directory exists
|
||||
os.makedirs(output_path, exist_ok=True)
|
||||
time.sleep(10)
|
||||
self._initialize_models() # Ensure models are initialized
|
||||
|
||||
def _initialize_models(self):
|
||||
# Set VSA environment variable
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
# os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
|
||||
# Import only when needed - use direct imports to avoid module-level execution
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
# Initialize T2V model
|
||||
if self.t2v_generator is None:
|
||||
print(f"Initializing T2V model: {self.t2v_model_path}")
|
||||
self.t2v_generator = VideoGenerator.from_pretrained(
|
||||
model_path=self.t2v_model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
# Adjust these offload parameters if you have < 32GB of VRAM
|
||||
text_encoder_cpu_offload=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
VSA_sparsity=0.8,
|
||||
# master_port=port,
|
||||
)
|
||||
self.t2v_default_params = SamplingParam.from_pretrained(self.t2v_model_path)
|
||||
print("✅ T2V model initialized successfully")
|
||||
|
||||
# Initialize I2V model
|
||||
# if self.i2v_generator is None:
|
||||
if False:
|
||||
print(f"Initializing I2V model: {self.i2v_model_path}")
|
||||
self.i2v_generator = VideoGenerator.from_pretrained(
|
||||
model_path=self.i2v_model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
# Adjust these offload parameters if you have < 32GB of VRAM
|
||||
text_encoder_cpu_offload=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
VSA_sparsity=0.8,
|
||||
# master_port=port,
|
||||
)
|
||||
self.i2v_default_params = SamplingParam.from_pretrained(self.i2v_model_path)
|
||||
print("✅ I2V model initialized successfully")
|
||||
|
||||
@app.post("/generate_video", response_model=VideoGenerationResponse)
|
||||
@limiter.limit("50/minute") # Allow 2 requests per minute per IP
|
||||
async def generate_video(self, request: Request, video_request: VideoGenerationRequest) -> VideoGenerationResponse:
|
||||
try:
|
||||
# Select the appropriate model and parameters based on model_type
|
||||
if video_request.model_type.lower() == "i2v":
|
||||
generator = self.i2v_generator
|
||||
params = deepcopy(self.i2v_default_params)
|
||||
print(f"Using I2V model for generation")
|
||||
else:
|
||||
generator = self.t2v_generator
|
||||
params = deepcopy(self.t2v_default_params)
|
||||
print(f"Using T2V model for generation")
|
||||
|
||||
# Update parameters with request values
|
||||
params.prompt = video_request.prompt
|
||||
# Only override negative prompt if user explicitly opts in
|
||||
if video_request.use_negative_prompt:
|
||||
params.negative_prompt = video_request.negative_prompt
|
||||
|
||||
params.seed = video_request.seed
|
||||
params.guidance_scale = video_request.guidance_scale
|
||||
params.num_frames = video_request.num_frames
|
||||
params.height = video_request.height
|
||||
params.width = video_request.width
|
||||
params.num_inference_steps = video_request.num_inference_steps
|
||||
|
||||
# Handle seed randomization
|
||||
if video_request.randomize_seed:
|
||||
params.seed = torch.randint(0, 1000000, (1,)).item()
|
||||
|
||||
# Ensure negative_prompt is a non-None string; FastVideo validation disallows None
|
||||
if params.negative_prompt is None:
|
||||
params.negative_prompt = "" # empty string satisfies validator
|
||||
|
||||
# Set up output path and video saving
|
||||
params.save_video = True
|
||||
params.output_path = self.output_path
|
||||
# params.return_frames = False # avoid keeping frames in memory
|
||||
|
||||
# Create a clean filename from the prompt
|
||||
safe_prompt = video_request.prompt[:100].replace(' ', '_').replace('/', '_').replace('\\', '_')
|
||||
|
||||
# Store desired video name inside the SamplingParam to avoid unknown kwarg errors
|
||||
setattr(params, "output_video_name", safe_prompt)
|
||||
|
||||
# Handle image_path for I2V
|
||||
if video_request.image_path:
|
||||
params.image_path = video_request.image_path
|
||||
|
||||
# Generate the video with proper output path and filename
|
||||
result = generator.generate_video(
|
||||
prompt=video_request.prompt,
|
||||
sampling_param=params,
|
||||
save_video=True, # Match the params.save_video setting
|
||||
)
|
||||
|
||||
# The actual output path where the video was saved
|
||||
output_path = os.path.join(self.output_path, f"{safe_prompt}.mp4")
|
||||
|
||||
# Verify the file exists
|
||||
if not os.path.exists(output_path):
|
||||
raise FileNotFoundError(f"Video was not saved to expected location: {output_path}")
|
||||
|
||||
frames = result.get("frames", [])
|
||||
|
||||
# Encode frames to base64 for web transmission only if requested
|
||||
encoded_frames = None
|
||||
if video_request.return_frames and frames:
|
||||
try:
|
||||
encoded_frames = encode_frames_to_base64(frames)
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to encode frames: {e}")
|
||||
encoded_frames = None
|
||||
|
||||
response = VideoGenerationResponse(
|
||||
output_path=output_path,
|
||||
frames=encoded_frames,
|
||||
seed=params.seed,
|
||||
success=True
|
||||
)
|
||||
|
||||
# Memory cleanup to avoid OOM in repeated generations
|
||||
import gc
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
return VideoGenerationResponse(
|
||||
output_path="",
|
||||
seed=video_request.seed,
|
||||
success=False,
|
||||
error_message=str(e)
|
||||
)
|
||||
|
||||
@app.get("/health")
|
||||
@limiter.limit("10/minute") # Allow 10 health checks per minute per IP
|
||||
async def health_check(self, request: Request):
|
||||
return {"status": "healthy"}
|
||||
|
||||
|
||||
def start_ray_serve(
|
||||
t2v_model_path: str = "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
i2v_model_path: str = "Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
output_path: str = "outputs",
|
||||
host: str = "0.0.0.0",
|
||||
port: int = 8000
|
||||
):
|
||||
"""Start the Ray Serve backend"""
|
||||
# Initialize Ray
|
||||
if not ray.is_initialized():
|
||||
ray.init()
|
||||
|
||||
# Deploy the API
|
||||
api = FastVideoAPI.bind(t2v_model_path, i2v_model_path, output_path)
|
||||
serve.run(api, route_prefix="/", name="fast_video") # detach
|
||||
|
||||
print(f"Ray Serve backend started at http://{host}:{port}")
|
||||
print(f"T2V Model: {t2v_model_path}")
|
||||
print(f"I2V Model: {i2v_model_path}")
|
||||
print(f"Health check: http://{host}:{port}/health")
|
||||
print(f"Video generation endpoint: http://{host}:{port}/generate_video")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description="FastVideo Ray Serve Backend")
|
||||
parser.add_argument("--t2v_model_path",
|
||||
type=str,
|
||||
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
help="Path to the T2V model")
|
||||
parser.add_argument("--i2v_model_path",
|
||||
type=str,
|
||||
default="Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
help="Path to the I2V model")
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs",
|
||||
help="Path to save generated videos")
|
||||
parser.add_argument("--host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Host to bind to")
|
||||
parser.add_argument("--port",
|
||||
type=int,
|
||||
default=8000,
|
||||
help="Port to bind to")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
start_ray_serve(
|
||||
t2v_model_path=args.t2v_model_path,
|
||||
i2v_model_path=args.i2v_model_path,
|
||||
output_path=args.output_path,
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
)
|
||||
|
||||
# ---- keep the process alive ---------------------------------
|
||||
import signal, sys, time
|
||||
signal.signal(signal.SIGINT, lambda *_: sys.exit(0)) # Ctrl-C
|
||||
signal.signal(signal.SIGTERM, lambda *_: sys.exit(0)) # docker stop etc.
|
||||
|
||||
print("✅ FastVideo backend is running. Press Ctrl-C to stop.")
|
||||
while True:
|
||||
time.sleep(3600)
|
||||
@@ -0,0 +1,334 @@
|
||||
import time
|
||||
import os
|
||||
import torch
|
||||
import base64
|
||||
import io
|
||||
from copy import deepcopy
|
||||
from typing import Dict, Any, Optional, List
|
||||
|
||||
import ray
|
||||
from ray import serve
|
||||
from fastapi import FastAPI, Request
|
||||
from pydantic import BaseModel
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
from slowapi import Limiter, _rate_limit_exceeded_handler
|
||||
from slowapi.util import get_remote_address
|
||||
from slowapi.errors import RateLimitExceeded
|
||||
|
||||
|
||||
class VideoGenerationRequest(BaseModel):
|
||||
prompt: str
|
||||
negative_prompt: Optional[str] = None
|
||||
use_negative_prompt: bool = False
|
||||
seed: int = 42
|
||||
guidance_scale: float = 7.5
|
||||
num_frames: int = 21
|
||||
height: int = 448
|
||||
width: int = 832
|
||||
num_inference_steps: int = 20
|
||||
randomize_seed: bool = False
|
||||
return_frames: bool = False # Whether to return base64 encoded frames
|
||||
image_path: Optional[str] = None # Path to input image for I2V
|
||||
model_type: str = "t2v" # "t2v" or "i2v" to specify which model to use
|
||||
|
||||
|
||||
class VideoGenerationResponse(BaseModel):
|
||||
output_path: str
|
||||
seed: int
|
||||
success: bool
|
||||
error_message: Optional[str] = None
|
||||
frames: Optional[List[str]] = None # Base64 encoded frames
|
||||
|
||||
|
||||
def encode_frames_to_base64(frames: List[np.ndarray]) -> List[str]:
|
||||
"""Convert numpy frames (0-255) to base64-encoded PNG images"""
|
||||
if not frames:
|
||||
return []
|
||||
|
||||
encoded_frames = []
|
||||
|
||||
for i, frame in enumerate(frames):
|
||||
try:
|
||||
# Ensure frame is numpy array
|
||||
if not isinstance(frame, np.ndarray):
|
||||
print(f"Warning: Frame {i} is not a numpy array, skipping")
|
||||
continue
|
||||
|
||||
# Ensure frame is uint8
|
||||
if frame.dtype != np.uint8:
|
||||
# Clip values to 0-255 range and convert to uint8
|
||||
frame = np.clip(frame, 0, 255).astype(np.uint8)
|
||||
|
||||
# Convert numpy array to PIL Image
|
||||
if len(frame.shape) == 3 and frame.shape[2] == 3:
|
||||
# RGB image
|
||||
pil_image = Image.fromarray(frame, mode='RGB')
|
||||
elif len(frame.shape) == 3 and frame.shape[2] == 4:
|
||||
# RGBA image
|
||||
pil_image = Image.fromarray(frame, mode='RGBA')
|
||||
elif len(frame.shape) == 2:
|
||||
# Grayscale image
|
||||
pil_image = Image.fromarray(frame, mode='L')
|
||||
else:
|
||||
print(f"Warning: Frame {i} has unsupported shape {frame.shape}, skipping")
|
||||
continue
|
||||
|
||||
# Save to bytes buffer as PNG
|
||||
buffer = io.BytesIO()
|
||||
pil_image.save(buffer, format='PNG')
|
||||
buffer.seek(0)
|
||||
|
||||
# Encode to base64
|
||||
img_base64 = base64.b64encode(buffer.getvalue()).decode('utf-8')
|
||||
encoded_frames.append(f"data:image/png;base64,{img_base64}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to encode frame {i}: {e}")
|
||||
continue
|
||||
|
||||
return encoded_frames
|
||||
|
||||
|
||||
# Create FastAPI app with rate limiting
|
||||
app = FastAPI()
|
||||
|
||||
# Initialize rate limiter
|
||||
limiter = Limiter(key_func=get_remote_address)
|
||||
app.state.limiter = limiter
|
||||
app.add_exception_handler(RateLimitExceeded, _rate_limit_exceeded_handler)
|
||||
|
||||
|
||||
@serve.deployment(
|
||||
num_replicas=3, # Set to 3 for cluster 0
|
||||
ray_actor_options={
|
||||
"num_cpus": 10,
|
||||
"num_gpus": 1,
|
||||
"runtime_env": {"conda": "fv"},
|
||||
},
|
||||
)
|
||||
@serve.ingress(app)
|
||||
class FastVideoMultiGPUAPI:
|
||||
def __init__(self, t2v_model_path: str, i2v_model_path: str, output_path: str, gpu_id: int = 0):
|
||||
self.t2v_model_path = t2v_model_path
|
||||
self.i2v_model_path = i2v_model_path
|
||||
self.output_path = output_path
|
||||
self.gpu_id = gpu_id
|
||||
|
||||
# Initialize the video generators
|
||||
self.t2v_generator = None
|
||||
self.i2v_generator = None
|
||||
self.t2v_default_params = None
|
||||
self.i2v_default_params = None
|
||||
|
||||
# Ensure output directory exists
|
||||
os.makedirs(output_path, exist_ok=True)
|
||||
time.sleep(10)
|
||||
self._initialize_models()
|
||||
|
||||
def _initialize_models(self):
|
||||
# Set VSA environment variable
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "FLASH_ATTN"
|
||||
|
||||
# Import only when needed
|
||||
from fastvideo.entrypoints.video_generator import VideoGenerator
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
# Initialize T2V model
|
||||
if False: # Disabled for now
|
||||
print(f"Initializing T2V model on GPU {self.gpu_id}: {self.t2v_model_path}")
|
||||
self.t2v_generator = VideoGenerator.from_pretrained(
|
||||
model_path=self.t2v_model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
text_encoder_cpu_offload=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
VSA_sparsity=0.8,
|
||||
)
|
||||
self.t2v_default_params = SamplingParam.from_pretrained(self.t2v_model_path)
|
||||
print(f"✅ T2V model initialized successfully on GPU {self.gpu_id}")
|
||||
|
||||
# Initialize I2V model
|
||||
if self.i2v_generator is None:
|
||||
print(f"Initializing I2V model on GPU {self.gpu_id}: {self.i2v_model_path}")
|
||||
self.i2v_generator = VideoGenerator.from_pretrained(
|
||||
model_path=self.i2v_model_path,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=True,
|
||||
text_encoder_cpu_offload=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
VSA_sparsity=0.8,
|
||||
)
|
||||
self.i2v_default_params = SamplingParam.from_pretrained(self.i2v_model_path)
|
||||
print(f"✅ I2V model initialized successfully on GPU {self.gpu_id}")
|
||||
|
||||
@app.post("/generate_video", response_model=VideoGenerationResponse)
|
||||
@limiter.limit("2/minute") # Allow 2 requests per minute per IP
|
||||
async def generate_video(self, request: Request, video_request: VideoGenerationRequest) -> VideoGenerationResponse:
|
||||
try:
|
||||
# Select the appropriate model and parameters based on model_type
|
||||
if video_request.model_type.lower() == "i2v":
|
||||
generator = self.i2v_generator
|
||||
params = deepcopy(self.i2v_default_params)
|
||||
print(f"Using I2V model for generation on GPU {self.gpu_id}")
|
||||
else:
|
||||
generator = self.t2v_generator
|
||||
params = deepcopy(self.t2v_default_params)
|
||||
print(f"Using T2V model for generation on GPU {self.gpu_id}")
|
||||
|
||||
# Update parameters with request values
|
||||
params.prompt = video_request.prompt
|
||||
|
||||
# Handle seed randomization
|
||||
if video_request.randomize_seed:
|
||||
params.seed = torch.randint(0, 1000000, (1,)).item()
|
||||
|
||||
# Ensure negative_prompt is a non-None string
|
||||
if params.negative_prompt is None:
|
||||
params.negative_prompt = ""
|
||||
|
||||
# Set up output path and video saving
|
||||
params.save_video = True
|
||||
params.output_path = self.output_path
|
||||
|
||||
# Create a clean filename from the prompt
|
||||
safe_prompt = video_request.prompt[:100].replace(' ', '_').replace('/', '_').replace('\\', '_')
|
||||
setattr(params, "output_video_name", safe_prompt)
|
||||
|
||||
# Handle image_path for I2V
|
||||
if video_request.image_path:
|
||||
params.image_path = video_request.image_path
|
||||
|
||||
# Generate the video
|
||||
result = generator.generate_video(
|
||||
prompt=video_request.prompt,
|
||||
sampling_param=params,
|
||||
save_video=True,
|
||||
)
|
||||
|
||||
frames = result.get("frames", [])
|
||||
|
||||
# Encode frames to base64 for web transmission only if requested
|
||||
encoded_frames = None
|
||||
if video_request.return_frames and frames:
|
||||
try:
|
||||
encoded_frames = encode_frames_to_base64(frames)
|
||||
except Exception as e:
|
||||
print(f"Warning: Failed to encode frames: {e}")
|
||||
encoded_frames = None
|
||||
|
||||
response = VideoGenerationResponse(
|
||||
output_path="",
|
||||
frames=encoded_frames,
|
||||
seed=params.seed,
|
||||
success=True
|
||||
)
|
||||
|
||||
# Memory cleanup to avoid OOM in repeated generations
|
||||
import gc
|
||||
gc.collect()
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
return VideoGenerationResponse(
|
||||
output_path="",
|
||||
seed=video_request.seed,
|
||||
success=False,
|
||||
error_message=str(e)
|
||||
)
|
||||
|
||||
@app.get("/health")
|
||||
@limiter.limit("10/minute") # Allow 10 health checks per minute per IP
|
||||
async def health_check(self, request: Request):
|
||||
return {"status": "healthy", "gpu_id": self.gpu_id}
|
||||
|
||||
|
||||
def start_ray_serve_multi_gpu(
|
||||
t2v_model_path: str = "FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
i2v_model_path: str = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
output_path: str = "outputs",
|
||||
host: str = "0.0.0.0",
|
||||
port: int = 8000,
|
||||
num_gpus: int = 8,
|
||||
cluster_id: int = 0
|
||||
):
|
||||
"""Start the Ray Serve backend with multiple GPU replicas"""
|
||||
# Initialize Ray
|
||||
if not ray.is_initialized():
|
||||
ray.init()
|
||||
|
||||
# Use unique application name based on cluster_id
|
||||
app_name = f"fast_video_cluster_{cluster_id}"
|
||||
|
||||
# Deploy the API
|
||||
api = FastVideoMultiGPUAPI.bind(t2v_model_path, i2v_model_path, output_path)
|
||||
serve.run(api, route_prefix=f"/cluster_{cluster_id}", name=app_name)
|
||||
|
||||
print(f"Ray Serve multi-GPU backend started at http://{host}:{port}")
|
||||
print(f"T2V Model: {t2v_model_path}")
|
||||
print(f"I2V Model: {i2v_model_path}")
|
||||
print(f"Number of GPU replicas: {num_gpus}")
|
||||
print(f"Cluster ID: {cluster_id}")
|
||||
print(f"Application name: {app_name}")
|
||||
print(f"Health check: http://{host}:{port}/cluster_{cluster_id}/health")
|
||||
print(f"Video generation endpoint: http://{host}:{port}/cluster_{cluster_id}/generate_video")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(description="FastVideo Ray Serve Multi-GPU Backend")
|
||||
parser.add_argument("--t2v_model_path",
|
||||
type=str,
|
||||
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
help="Path to the T2V model")
|
||||
parser.add_argument("--i2v_model_path",
|
||||
type=str,
|
||||
default="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
help="Path to the I2V model")
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs",
|
||||
help="Path to save generated videos")
|
||||
parser.add_argument("--host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Host to bind to")
|
||||
parser.add_argument("--port",
|
||||
type=int,
|
||||
default=8000,
|
||||
help="Port to bind to")
|
||||
parser.add_argument("--num_gpus",
|
||||
type=int,
|
||||
default=8,
|
||||
help="Number of GPU replicas")
|
||||
parser.add_argument("--cluster_id",
|
||||
type=int,
|
||||
default=0,
|
||||
help="Cluster ID for unique naming")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
start_ray_serve_multi_gpu(
|
||||
t2v_model_path=args.t2v_model_path,
|
||||
i2v_model_path=args.i2v_model_path,
|
||||
output_path=args.output_path,
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
num_gpus=args.num_gpus,
|
||||
cluster_id=args.cluster_id,
|
||||
)
|
||||
|
||||
# Keep the process alive
|
||||
import signal, sys, time
|
||||
signal.signal(signal.SIGINT, lambda *_: sys.exit(0))
|
||||
signal.signal(signal.SIGTERM, lambda *_: sys.exit(0))
|
||||
|
||||
print("✅ FastVideo multi-GPU backend is running. Press Ctrl-C to stop.")
|
||||
while True:
|
||||
time.sleep(3600)
|
||||
@@ -0,0 +1,232 @@
|
||||
"""
|
||||
Startup script for FastVideo with Ray Serve backend and Gradio frontend.
|
||||
This script starts both the backend and frontend services.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
import threading
|
||||
import signal
|
||||
import requests
|
||||
from pathlib import Path
|
||||
|
||||
# Add the project root to the Python path
|
||||
project_root = Path(__file__).parent.parent.parent.parent
|
||||
sys.path.insert(0, str(project_root))
|
||||
|
||||
|
||||
def check_backend_health(backend_url: str, max_retries: int = 100) -> bool:
|
||||
"""Check if the backend is healthy"""
|
||||
health_url = f"{backend_url}/health"
|
||||
|
||||
for i in range(max_retries):
|
||||
try:
|
||||
response = requests.get(health_url, timeout=5)
|
||||
if response.status_code == 200:
|
||||
print(f"✅ Backend is healthy at {backend_url}")
|
||||
return True
|
||||
except requests.exceptions.RequestException:
|
||||
pass
|
||||
|
||||
if i < max_retries - 1:
|
||||
print(f"⏳ Waiting for backend to start... ({i+1}/{max_retries})")
|
||||
time.sleep(2)
|
||||
|
||||
print(f"❌ Backend failed to start within {max_retries * 2} seconds")
|
||||
return False
|
||||
|
||||
|
||||
def start_backend(args):
|
||||
"""Start the Ray Serve backend"""
|
||||
backend_script = Path(__file__).parent / "ray_serve_backend.py"
|
||||
|
||||
cmd = [
|
||||
sys.executable, str(backend_script),
|
||||
"--t2v_model_path", args.t2v_model_path,
|
||||
"--i2v_model_path", args.i2v_model_path,
|
||||
"--output_path", args.output_path,
|
||||
"--host", args.backend_host,
|
||||
"--port", str(args.backend_port)
|
||||
]
|
||||
|
||||
print(f"🚀 Starting Ray Serve backend...")
|
||||
print(f"Command: {' '.join(cmd)}")
|
||||
|
||||
# Start the backend process
|
||||
backend_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1
|
||||
)
|
||||
|
||||
# Monitor backend output
|
||||
def monitor_backend():
|
||||
for line in backend_process.stdout:
|
||||
print(f"[BACKEND] {line.rstrip()}")
|
||||
|
||||
monitor_thread = threading.Thread(target=monitor_backend, daemon=True)
|
||||
monitor_thread.start()
|
||||
|
||||
return backend_process
|
||||
|
||||
|
||||
def start_frontend(args):
|
||||
"""Start the Gradio frontend"""
|
||||
frontend_script = Path(__file__).parent / "gradio_frontend.py"
|
||||
backend_url = f"http://{args.backend_host}:{args.backend_port}"
|
||||
|
||||
cmd = [
|
||||
sys.executable, str(frontend_script),
|
||||
"--backend_url", backend_url,
|
||||
"--t2v_model_path", args.t2v_model_path,
|
||||
"--i2v_model_path", args.i2v_model_path,
|
||||
"--host", args.frontend_host,
|
||||
"--port", str(args.frontend_port)
|
||||
]
|
||||
|
||||
print(f"🎨 Starting Gradio frontend...")
|
||||
print(f"Command: {' '.join(cmd)}")
|
||||
|
||||
# Start the frontend process
|
||||
frontend_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1
|
||||
)
|
||||
|
||||
# Monitor frontend output
|
||||
def monitor_frontend():
|
||||
for line in frontend_process.stdout:
|
||||
print(f"[FRONTEND] {line.rstrip()}")
|
||||
|
||||
monitor_thread = threading.Thread(target=monitor_frontend, daemon=True)
|
||||
monitor_thread.start()
|
||||
|
||||
return frontend_process
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="FastVideo Ray Serve App")
|
||||
|
||||
# Model and output settings
|
||||
parser.add_argument("--t2v_model_path",
|
||||
type=str,
|
||||
default="Wan-AI/Wan2.2-TI2V-5BDiffusers",
|
||||
help="Path to the T2V model")
|
||||
parser.add_argument("--i2v_model_path",
|
||||
type=str,
|
||||
default="Wan-AI/Wan2.2-TI2V-5B-Diffusers",
|
||||
help="Path to the I2V model")
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs",
|
||||
help="Path to save generated videos")
|
||||
|
||||
# Backend settings
|
||||
parser.add_argument("--backend_host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Backend host to bind to")
|
||||
parser.add_argument("--backend_port",
|
||||
type=int,
|
||||
default=8000,
|
||||
help="Backend port to bind to")
|
||||
|
||||
# Frontend settings
|
||||
parser.add_argument("--frontend_host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Frontend host to bind to")
|
||||
parser.add_argument("--frontend_port",
|
||||
type=int,
|
||||
default=7860,
|
||||
help="Frontend port to bind to")
|
||||
|
||||
# Other settings
|
||||
parser.add_argument("--skip_backend_check",
|
||||
action="store_true",
|
||||
help="Skip backend health check")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Ensure output directory exists
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
|
||||
print("🎬 FastVideo Ray Serve App")
|
||||
print("=" * 50)
|
||||
print(f"T2V Model: {args.t2v_model_path}")
|
||||
print(f"I2V Model: {args.i2v_model_path}")
|
||||
print(f"Output: {args.output_path}")
|
||||
print(f"Backend: http://{args.backend_host}:{args.backend_port}")
|
||||
print(f"Frontend: http://{args.frontend_host}:{args.frontend_port}")
|
||||
print("=" * 50)
|
||||
|
||||
# Start backend
|
||||
backend_process = start_backend(args)
|
||||
|
||||
# Wait for backend to be ready
|
||||
backend_url = f"http://{args.backend_host}:{args.backend_port}"
|
||||
|
||||
if not args.skip_backend_check:
|
||||
if not check_backend_health(backend_url):
|
||||
print("❌ Backend failed to start. Terminating...")
|
||||
backend_process.terminate()
|
||||
sys.exit(1)
|
||||
|
||||
# Start frontend
|
||||
frontend_process = start_frontend(args)
|
||||
|
||||
print("\n🎉 Both services are starting up!")
|
||||
print(f"📺 Frontend will be available at: http://{args.frontend_host}:{args.frontend_port}")
|
||||
print(f"🔧 Backend API will be available at: {backend_url}")
|
||||
print("\nPress Ctrl+C to stop both services...")
|
||||
# return
|
||||
|
||||
# Signal handler for graceful shutdown
|
||||
def signal_handler(signum, frame):
|
||||
print("\n🛑 Shutting down services...")
|
||||
frontend_process.terminate()
|
||||
backend_process.terminate()
|
||||
|
||||
# Wait for processes to terminate
|
||||
try:
|
||||
frontend_process.wait(timeout=5)
|
||||
backend_process.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
print("⚠️ Force killing processes...")
|
||||
frontend_process.kill()
|
||||
backend_process.kill()
|
||||
|
||||
print("✅ Services stopped")
|
||||
sys.exit(0)
|
||||
|
||||
signal.signal(signal.SIGINT, signal_handler)
|
||||
signal.signal(signal.SIGTERM, signal_handler)
|
||||
|
||||
# Monitor processes
|
||||
try:
|
||||
while True:
|
||||
# Check if processes are still running
|
||||
if frontend_process.poll() is not None:
|
||||
print("❌ Frontend process died unexpectedly")
|
||||
break
|
||||
|
||||
if backend_process.poll() is not None:
|
||||
print("❌ Backend process died unexpectedly")
|
||||
break
|
||||
|
||||
time.sleep(1)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
signal_handler(signal.SIGINT, None)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,465 @@
|
||||
"""
|
||||
Startup script for FastVideo scalable architecture:
|
||||
ngrok -> nginx reverse proxy -> frontend1/frontend2 -> backend1×8/backend2×8
|
||||
|
||||
This script starts:
|
||||
1. Multiple backend instances (8 GPU replicas each)
|
||||
2. Multiple frontend instances (2 instances)
|
||||
3. Nginx reverse proxy
|
||||
4. Optional ngrok tunnel
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
import threading
|
||||
import signal
|
||||
import requests
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
# Add the project root to the Python path
|
||||
project_root = Path(__file__).parent.parent.parent.parent
|
||||
sys.path.insert(0, str(project_root))
|
||||
|
||||
|
||||
def check_service_health(url: str, max_retries: int = 50) -> bool:
|
||||
"""Check if a service is healthy"""
|
||||
for i in range(max_retries):
|
||||
try:
|
||||
response = requests.get(url, timeout=5)
|
||||
if response.status_code == 200:
|
||||
print(f"✅ Service is healthy at {url}")
|
||||
return True
|
||||
except requests.exceptions.RequestException:
|
||||
pass
|
||||
|
||||
if i < max_retries - 1:
|
||||
print(f"⏳ Waiting for service to start... ({i+1}/{max_retries})")
|
||||
time.sleep(2)
|
||||
|
||||
print(f"❌ Service failed to start within {max_retries * 2} seconds")
|
||||
return False
|
||||
|
||||
|
||||
def start_backend_cluster(args, cluster_id: int):
|
||||
"""Start one backend cluster (Ray-Serve application)."""
|
||||
backend_script = Path(__file__).parent / "ray_serve_backend_scalable.py"
|
||||
|
||||
# All Ray Serve apps share the same HTTP server (default 8000).
|
||||
# We still forward the port flag for completeness, but keep it
|
||||
# identical for every cluster.
|
||||
base_port = args.backend_base_port
|
||||
|
||||
cmd = [
|
||||
sys.executable, str(backend_script),
|
||||
"--t2v_model_path", args.t2v_model_path,
|
||||
"--i2v_model_path", args.i2v_model_path,
|
||||
"--output_path", args.output_path,
|
||||
"--host", args.backend_host,
|
||||
"--port", str(base_port),
|
||||
"--num_gpus", str(args.num_gpus_per_cluster),
|
||||
"--cluster_id", str(cluster_id),
|
||||
]
|
||||
|
||||
print(f"🚀 Starting Backend Cluster {cluster_id + 1} (HTTP port {base_port})...")
|
||||
print(f"Command: {' '.join(cmd)}")
|
||||
|
||||
# Start the backend process
|
||||
backend_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1
|
||||
)
|
||||
|
||||
# Monitor backend output
|
||||
def monitor_backend():
|
||||
for line in backend_process.stdout:
|
||||
print(f"[BACKEND-{cluster_id + 1}] {line.rstrip()}")
|
||||
|
||||
monitor_thread = threading.Thread(target=monitor_backend, daemon=True)
|
||||
monitor_thread.start()
|
||||
|
||||
return backend_process, base_port
|
||||
|
||||
|
||||
def start_frontend_instance(args, instance_id: int, backend_url: str):
|
||||
"""Start a single frontend instance"""
|
||||
frontend_script = Path(__file__).parent / "gradio_frontend.py"
|
||||
frontend_port = args.frontend_base_port + instance_id
|
||||
|
||||
# Update backend URL to include cluster-specific path
|
||||
cluster_id = instance_id % args.num_backend_clusters
|
||||
backend_url_with_cluster = f"{backend_url}/cluster_{cluster_id}"
|
||||
|
||||
cmd = [
|
||||
sys.executable, str(frontend_script),
|
||||
"--backend_url", backend_url_with_cluster,
|
||||
"--t2v_model_path", args.t2v_model_path,
|
||||
"--i2v_model_path", args.i2v_model_path,
|
||||
"--host", args.frontend_host,
|
||||
"--port", str(frontend_port)
|
||||
]
|
||||
|
||||
print(f"🎨 Starting Frontend {instance_id + 1} on port {frontend_port}...")
|
||||
print(f"Command: {' '.join(cmd)}")
|
||||
|
||||
# Start the frontend process
|
||||
frontend_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1
|
||||
)
|
||||
|
||||
# Monitor frontend output
|
||||
def monitor_frontend():
|
||||
for line in frontend_process.stdout:
|
||||
print(f"[FRONTEND-{instance_id + 1}] {line.rstrip()}")
|
||||
|
||||
monitor_thread = threading.Thread(target=monitor_frontend, daemon=True)
|
||||
monitor_thread.start()
|
||||
|
||||
return frontend_process, frontend_port
|
||||
|
||||
|
||||
def start_nginx(args):
|
||||
"""Start nginx reverse proxy"""
|
||||
nginx_conf = Path(__file__).parent / "nginx.conf"
|
||||
|
||||
# Update nginx configuration with actual ports
|
||||
update_nginx_config(args)
|
||||
|
||||
cmd = [
|
||||
"nginx",
|
||||
"-c", str(nginx_conf),
|
||||
"-g", "daemon off;"
|
||||
]
|
||||
|
||||
print(f"🌐 Starting Nginx reverse proxy...")
|
||||
print(f"Command: {' '.join(cmd)}")
|
||||
|
||||
# Start nginx process
|
||||
nginx_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1
|
||||
)
|
||||
|
||||
# Monitor nginx output
|
||||
def monitor_nginx():
|
||||
for line in nginx_process.stdout:
|
||||
print(f"[NGINX] {line.rstrip()}")
|
||||
|
||||
monitor_thread = threading.Thread(target=monitor_nginx, daemon=True)
|
||||
monitor_thread.start()
|
||||
|
||||
return nginx_process
|
||||
|
||||
|
||||
def update_nginx_config(args):
|
||||
"""Rewrite nginx.conf with the correct ports – NO “/cluster_X” in upstreams."""
|
||||
nginx_conf = Path(__file__).parent / "nginx.conf"
|
||||
nginx_conf_backup = Path(__file__).parent / "nginx.conf.backup"
|
||||
|
||||
if not nginx_conf_backup.exists():
|
||||
nginx_conf_backup.write_text(nginx_conf.read_text())
|
||||
|
||||
config_content = nginx_conf_backup.read_text()
|
||||
|
||||
# ── 1. front-end pool ───────────────────────────────────────────────
|
||||
frontend_servers = "\n ".join(
|
||||
f"server 127.0.0.1:{args.frontend_base_port + i} "
|
||||
f"weight=1 max_fails=3 fail_timeout=30s;"
|
||||
for i in range(args.num_frontends)
|
||||
)
|
||||
config_content = config_content.replace(
|
||||
"# Upstream for frontend load balancing",
|
||||
f"# Upstream for frontend load balancing\n upstream frontend_servers {{\n"
|
||||
f" # Round-robin load balancing between frontends\n {frontend_servers}"
|
||||
)
|
||||
|
||||
# Shared Ray-Serve HTTP port
|
||||
backend_port = args.backend_base_port # default 8000
|
||||
backend_line = (f"server 127.0.0.1:{backend_port} "
|
||||
f"weight=1 max_fails=3 fail_timeout=30s;")
|
||||
|
||||
# ── 2. backend-1 pool ───────────────────────────────────────────────
|
||||
config_content = config_content.replace(
|
||||
"# Upstream for backend1 load balancing",
|
||||
f"# Upstream for backend1 load balancing\n upstream backend1_servers {{\n"
|
||||
f" {backend_line}"
|
||||
)
|
||||
|
||||
# ── 3. backend-2 pool ───────────────────────────────────────────────
|
||||
config_content = config_content.replace(
|
||||
"# Upstream for backend2 load balancing",
|
||||
f"# Upstream for backend2 load balancing\n upstream backend2_servers {{\n"
|
||||
f" {backend_line}"
|
||||
)
|
||||
|
||||
# ── 4. strip any stray “/cluster_X” fragments ───────────────────────
|
||||
config_content = config_content.replace("/cluster_0", "").replace("/cluster_1", "")
|
||||
|
||||
# ── 5. use user-writable log directory ---------------------------------
|
||||
log_dir = Path(args.output_path).resolve()
|
||||
config_content = config_content.replace(
|
||||
"access_log /var/log/nginx/access.log;",
|
||||
f"access_log {log_dir}/nginx_access.log;")
|
||||
config_content = config_content.replace(
|
||||
"error_log /var/log/nginx/error.log;",
|
||||
f"error_log {log_dir}/nginx_error.log;")
|
||||
|
||||
nginx_conf.write_text(config_content)
|
||||
print("✅ nginx.conf updated (no path suffixes & custom log paths)")
|
||||
|
||||
|
||||
def start_ngrok(args):
|
||||
"""Start ngrok tunnel"""
|
||||
if not args.use_ngrok:
|
||||
return None
|
||||
|
||||
cmd = [
|
||||
"ngrok",
|
||||
"http",
|
||||
str(args.nginx_port),
|
||||
"--log=stdout"
|
||||
]
|
||||
|
||||
print(f"🌍 Starting ngrok tunnel to port {args.nginx_port}...")
|
||||
print(f"Command: {' '.join(cmd)}")
|
||||
|
||||
# Start ngrok process
|
||||
ngrok_process = subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.STDOUT,
|
||||
universal_newlines=True,
|
||||
bufsize=1
|
||||
)
|
||||
|
||||
# Monitor ngrok output
|
||||
def monitor_ngrok():
|
||||
for line in ngrok_process.stdout:
|
||||
print(f"[NGROK] {line.rstrip()}")
|
||||
|
||||
monitor_thread = threading.Thread(target=monitor_ngrok, daemon=True)
|
||||
monitor_thread.start()
|
||||
|
||||
return ngrok_process
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="FastVideo Scalable Architecture Launcher")
|
||||
|
||||
# Model and output settings
|
||||
parser.add_argument("--t2v_model_path",
|
||||
type=str,
|
||||
default="FastVideo/FastWan2.1-T2V-1.3B-Diffusers",
|
||||
help="Path to the T2V model")
|
||||
parser.add_argument("--i2v_model_path",
|
||||
type=str,
|
||||
default="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers",
|
||||
help="Path to the I2V model")
|
||||
parser.add_argument("--output_path",
|
||||
type=str,
|
||||
default="outputs",
|
||||
help="Path to save generated videos")
|
||||
|
||||
# Backend settings
|
||||
parser.add_argument("--backend_host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Backend host to bind to")
|
||||
parser.add_argument("--backend_base_port",
|
||||
type=int,
|
||||
default=8000,
|
||||
help="Base port for backend clusters")
|
||||
parser.add_argument("--num_backend_clusters",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Number of backend clusters")
|
||||
parser.add_argument("--num_gpus_per_cluster",
|
||||
type=int,
|
||||
default=3, # Changed from 8 to 3 (3+3=6 GPUs total, leaving 1 GPU buffer)
|
||||
help="Number of GPUs per backend cluster")
|
||||
|
||||
# Frontend settings
|
||||
parser.add_argument("--frontend_host",
|
||||
type=str,
|
||||
default="0.0.0.0",
|
||||
help="Frontend host to bind to")
|
||||
parser.add_argument("--frontend_base_port",
|
||||
type=int,
|
||||
default=7860,
|
||||
help="Base port for frontend instances")
|
||||
parser.add_argument("--num_frontends",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Number of frontend instances")
|
||||
|
||||
# Nginx settings
|
||||
parser.add_argument("--nginx_port",
|
||||
type=int,
|
||||
default=80,
|
||||
help="Port for nginx reverse proxy")
|
||||
|
||||
# Ngrok settings
|
||||
parser.add_argument("--use_ngrok",
|
||||
action="store_true",
|
||||
help="Start ngrok tunnel")
|
||||
|
||||
# Other settings
|
||||
parser.add_argument("--skip_health_check",
|
||||
action="store_true",
|
||||
help="Skip health checks")
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
# Ensure output directory exists
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
|
||||
print(" FastVideo Scalable Architecture")
|
||||
print("=" * 60)
|
||||
print(f"Architecture: ngrok -> nginx -> frontend1/frontend2 -> backend1×{args.num_gpus_per_cluster}/backend2×{args.num_gpus_per_cluster}")
|
||||
print(f"T2V Model: {args.t2v_model_path}")
|
||||
print(f"I2V Model: {args.i2v_model_path}")
|
||||
print(f"Output: {args.output_path}")
|
||||
print(f"Backend Clusters: {args.num_backend_clusters}")
|
||||
print(f"GPUs per Cluster: {args.num_gpus_per_cluster}")
|
||||
print(f"Total GPUs needed: {args.num_backend_clusters * args.num_gpus_per_cluster}")
|
||||
print(f"Frontend Instances: {args.num_frontends}")
|
||||
print(f"Nginx Port: {args.nginx_port}")
|
||||
print(f"Use Ngrok: {args.use_ngrok}")
|
||||
print("=" * 60)
|
||||
|
||||
# Start backend clusters
|
||||
backend_processes = []
|
||||
backend_urls = []
|
||||
|
||||
for i in range(args.num_backend_clusters):
|
||||
process, _ = start_backend_cluster(args, i)
|
||||
backend_processes.append(process)
|
||||
backend_urls.append(f"http://{args.backend_host}:{args.backend_base_port}")
|
||||
|
||||
# Wait for backends to be ready
|
||||
if not args.skip_health_check:
|
||||
print("\n⏳ Waiting for backend clusters to start...")
|
||||
for i, url in enumerate(backend_urls):
|
||||
if not check_service_health(f"{url}/cluster_{i}/health"):
|
||||
print(f"❌ Backend cluster {i + 1} failed to start. Terminating...")
|
||||
for process in backend_processes:
|
||||
process.terminate()
|
||||
sys.exit(1)
|
||||
|
||||
# Start frontend instances
|
||||
frontend_processes = []
|
||||
frontend_urls = []
|
||||
|
||||
for i in range(args.num_frontends):
|
||||
# Each frontend connects to a different backend cluster
|
||||
backend_url = backend_urls[i % len(backend_urls)]
|
||||
process, port = start_frontend_instance(args, i, backend_url)
|
||||
frontend_processes.append(process)
|
||||
frontend_urls.append(f"http://{args.frontend_host}:{port}")
|
||||
|
||||
# Wait for frontends to be ready
|
||||
if not args.skip_health_check:
|
||||
print("\n⏳ Waiting for frontend instances to start...")
|
||||
for i, url in enumerate(frontend_urls):
|
||||
if not check_service_health(url):
|
||||
print(f"❌ Frontend {i + 1} failed to start. Terminating...")
|
||||
for process in backend_processes + frontend_processes:
|
||||
process.terminate()
|
||||
sys.exit(1)
|
||||
|
||||
# Start nginx reverse proxy
|
||||
nginx_process = start_nginx(args)
|
||||
|
||||
# Wait for nginx to be ready
|
||||
if not args.skip_health_check:
|
||||
print("\n⏳ Waiting for nginx to start...")
|
||||
if not check_service_health(f"http://localhost:{args.nginx_port}/health"):
|
||||
print("❌ Nginx failed to start. Terminating...")
|
||||
for process in backend_processes + frontend_processes + [nginx_process]:
|
||||
process.terminate()
|
||||
sys.exit(1)
|
||||
|
||||
# Start ngrok tunnel (optional)
|
||||
ngrok_process = start_ngrok(args)
|
||||
|
||||
print("\n🎉 All services are starting up!")
|
||||
print(f"🌐 Nginx reverse proxy: http://localhost:{args.nginx_port}")
|
||||
for i, url in enumerate(frontend_urls):
|
||||
print(f"📺 Frontend {i + 1}: {url}")
|
||||
for i, url in enumerate(backend_urls):
|
||||
print(f" Backend Cluster {i + 1}: {url}")
|
||||
if args.use_ngrok:
|
||||
print("🌍 Ngrok tunnel is starting...")
|
||||
print("\nPress Ctrl+C to stop all services...")
|
||||
|
||||
# Signal handler for graceful shutdown
|
||||
def signal_handler(signum, frame):
|
||||
print("\n🛑 Shutting down all services...")
|
||||
all_processes = backend_processes + frontend_processes + [nginx_process]
|
||||
if ngrok_process:
|
||||
all_processes.append(ngrok_process)
|
||||
|
||||
for process in all_processes:
|
||||
if process:
|
||||
process.terminate()
|
||||
|
||||
# Wait for processes to terminate
|
||||
try:
|
||||
for process in all_processes:
|
||||
if process:
|
||||
process.wait(timeout=5)
|
||||
except subprocess.TimeoutExpired:
|
||||
print("⚠️ Force killing processes...")
|
||||
for process in all_processes:
|
||||
if process:
|
||||
process.kill()
|
||||
|
||||
print("✅ All services stopped")
|
||||
sys.exit(0)
|
||||
|
||||
signal.signal(signal.SIGINT, signal_handler)
|
||||
signal.signal(signal.SIGTERM, signal_handler)
|
||||
|
||||
# Monitor processes
|
||||
try:
|
||||
while True:
|
||||
# Check if processes are still running
|
||||
for i, process in enumerate(backend_processes):
|
||||
if process.poll() is not None:
|
||||
print(f"❌ Backend cluster {i + 1} process died unexpectedly")
|
||||
break
|
||||
|
||||
for i, process in enumerate(frontend_processes):
|
||||
if process.poll() is not None:
|
||||
print(f"❌ Frontend {i + 1} process died unexpectedly")
|
||||
break
|
||||
|
||||
if nginx_process and nginx_process.poll() is not None:
|
||||
print("❌ Nginx process died unexpectedly")
|
||||
break
|
||||
|
||||
if ngrok_process and ngrok_process.poll() is not None:
|
||||
print("❌ Ngrok process died unexpectedly")
|
||||
break
|
||||
|
||||
time.sleep(1)
|
||||
|
||||
except KeyboardInterrupt:
|
||||
signal_handler(signal.SIGINT, None)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,147 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Test script for T2V and I2V functionality in FastVideo Gradio app.
|
||||
This script tests both the backend and frontend modifications.
|
||||
"""
|
||||
|
||||
import requests
|
||||
import json
|
||||
import os
|
||||
from PIL import Image
|
||||
import numpy as np
|
||||
|
||||
def test_backend_t2v():
|
||||
"""Test the backend T2V functionality directly"""
|
||||
backend_url = "http://localhost:8000"
|
||||
|
||||
try:
|
||||
# Test T2V request data
|
||||
request_data = {
|
||||
"prompt": "A beautiful sunset over the ocean with gentle waves",
|
||||
"negative_prompt": "",
|
||||
"use_negative_prompt": False,
|
||||
"seed": 42,
|
||||
"guidance_scale": 7.5,
|
||||
"num_frames": 21,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_inference_steps": 20,
|
||||
"randomize_seed": False,
|
||||
"return_frames": True,
|
||||
"image_path": None,
|
||||
"model_type": "t2v"
|
||||
}
|
||||
|
||||
# Send request to backend
|
||||
response = requests.post(
|
||||
f"{backend_url}/generate_video",
|
||||
json=request_data,
|
||||
timeout=300 # 5 minutes timeout
|
||||
)
|
||||
|
||||
if response.status_code == 200:
|
||||
result = response.json()
|
||||
print("✅ Backend T2V test successful!")
|
||||
print(f"Success: {result.get('success')}")
|
||||
print(f"Seed used: {result.get('seed')}")
|
||||
if result.get('frames'):
|
||||
print(f"Frames returned: {len(result.get('frames'))}")
|
||||
else:
|
||||
print("No frames returned")
|
||||
else:
|
||||
print(f"❌ Backend T2V test failed with status {response.status_code}")
|
||||
print(f"Response: {response.text}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"❌ Backend T2V test failed with exception: {e}")
|
||||
|
||||
def test_backend_i2v():
|
||||
"""Test the backend I2V functionality directly"""
|
||||
backend_url = "http://localhost:8000"
|
||||
|
||||
# Create a simple test image
|
||||
test_image = Image.new('RGB', (256, 256), color='red')
|
||||
temp_image_path = "test_image.png"
|
||||
test_image.save(temp_image_path)
|
||||
|
||||
try:
|
||||
# Test I2V request data
|
||||
request_data = {
|
||||
"prompt": "The red square gently animates with subtle movement",
|
||||
"negative_prompt": "",
|
||||
"use_negative_prompt": False,
|
||||
"seed": 42,
|
||||
"guidance_scale": 7.5,
|
||||
"num_frames": 21,
|
||||
"height": 448,
|
||||
"width": 832,
|
||||
"num_inference_steps": 20,
|
||||
"randomize_seed": False,
|
||||
"return_frames": True,
|
||||
"image_path": temp_image_path,
|
||||
"model_type": "i2v"
|
||||
}
|
||||
|
||||
# Send request to backend
|
||||
response = requests.post(
|
||||
f"{backend_url}/generate_video",
|
||||
json=request_data,
|
||||
timeout=300 # 5 minutes timeout
|
||||
)
|
||||
|
||||
if response.status_code == 200:
|
||||
result = response.json()
|
||||
print("✅ Backend I2V test successful!")
|
||||
print(f"Success: {result.get('success')}")
|
||||
print(f"Seed used: {result.get('seed')}")
|
||||
if result.get('frames'):
|
||||
print(f"Frames returned: {len(result.get('frames'))}")
|
||||
else:
|
||||
print("No frames returned")
|
||||
else:
|
||||
print(f"❌ Backend I2V test failed with status {response.status_code}")
|
||||
print(f"Response: {response.text}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"❌ Backend I2V test failed with exception: {e}")
|
||||
|
||||
finally:
|
||||
# Clean up test image
|
||||
if os.path.exists(temp_image_path):
|
||||
os.remove(temp_image_path)
|
||||
|
||||
def test_backend_health():
|
||||
"""Test if the backend is running"""
|
||||
backend_url = "http://localhost:8000"
|
||||
|
||||
try:
|
||||
response = requests.get(f"{backend_url}/health", timeout=5)
|
||||
if response.status_code == 200:
|
||||
print("✅ Backend is healthy")
|
||||
return True
|
||||
else:
|
||||
print(f"❌ Backend health check failed: {response.status_code}")
|
||||
return False
|
||||
except Exception as e:
|
||||
print(f"❌ Backend health check failed: {e}")
|
||||
return False
|
||||
|
||||
if __name__ == "__main__":
|
||||
print("🧪 Testing FastVideo T2V and I2V functionality...")
|
||||
print("=" * 50)
|
||||
|
||||
# Test backend health first
|
||||
if test_backend_health():
|
||||
# Test T2V functionality
|
||||
print("\n📝 Testing T2V functionality...")
|
||||
test_backend_t2v()
|
||||
|
||||
# Test I2V functionality
|
||||
print("\n🖼️ Testing I2V functionality...")
|
||||
test_backend_i2v()
|
||||
else:
|
||||
print("⚠️ Backend is not running. Please start the backend first.")
|
||||
print("You can start it with: python start_ray_serve_app.py")
|
||||
|
||||
print("=" * 50)
|
||||
print("Test completed!")
|
||||
@@ -1,12 +1,16 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "./lora"
|
||||
OUTPUT_PATH = "./lora_out"
|
||||
def main():
|
||||
# Initialize VideoGenerator with the Wan model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=2,
|
||||
num_gpus=1,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
lora_path="benjamin-paine/steamboat-willie-1.3b",
|
||||
lora_nickname="steamboat"
|
||||
)
|
||||
@@ -16,6 +20,7 @@ def main():
|
||||
"num_frames": 81,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 32,
|
||||
"seed": 42,
|
||||
}
|
||||
# Generate video with LoRA style
|
||||
prompt = "steamboat willie style, golden era animation, close-up of a short fluffy monster kneeling beside a melting red candle. the mood is one of wonder and curiosity, as the monster gazes at the flame with wide eyes and open mouth. Its pose and expression convey a sense of innocence and playfulness, as if it is exploring the world around it for the first time. The use of warm colors and dramatic lighting further enhances the cozy atmosphere of the image."
|
||||
@@ -29,7 +34,7 @@ def main():
|
||||
negative_prompt=negative_prompt,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
|
||||
generator.set_lora_adapter(lora_nickname="flat_color", lora_path="motimalu/wan-flat-color-1.3b-v2")
|
||||
prompt = "flat color, no lineart, blending, negative space, artist:[john kafka|ponsuke kaikai|hara id 21|yoneyama mai|fuzichoco], 1girl, sakura miko, pink hair, cowboy shot, white shirt, floral print, off shoulder, outdoors, cherry blossom, tree shade, wariza, looking up, falling petals, half-closed eyes, white sky, clouds, live2d animation, upper body, high quality cinematic video of a woman sitting under a sakura tree. Dreamy and lonely, the camera close-ups on the face of the woman as she turns towards the viewer. The Camera is steady, This is a cowboy shot. The animation is smooth and fluid."
|
||||
negative_prompt = "bad quality video,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
|
||||
@@ -0,0 +1,39 @@
|
||||
"""
|
||||
Inference using a LoRA checkpoint from FastVideo trainer.
|
||||
"""
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "./lora_out"
|
||||
def main():
|
||||
# Initialize VideoGenerator with the Wan model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
lora_path="checkpoints/wan_t2v_finetune_lora/checkpoint-1250/transformer",
|
||||
lora_nickname="crush_smol"
|
||||
)
|
||||
kwargs = {
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 50,
|
||||
"seed": 42,
|
||||
}
|
||||
# Generate video with LoRA style
|
||||
prompt = "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table."
|
||||
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -11,6 +11,10 @@ def main():
|
||||
gen = VideoGenerator.from_pretrained(
|
||||
model_path="Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=1,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
)
|
||||
load_time = time.perf_counter() - start_time
|
||||
print(f"Model loading time: {load_time:.2f} seconds")
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user