Compare commits
44
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ea8c154176 | ||
|
|
580d6dfe1f | ||
|
|
344e43006a | ||
|
|
c5155b256e | ||
|
|
e005c7f3ac | ||
|
|
ff5a79ef60 | ||
|
|
ab01dc4ba5 | ||
|
|
285a950c1b | ||
|
|
46a0a85d85 | ||
|
|
4aeabbc629 | ||
|
|
949bb5c835 | ||
|
|
aab74c1271 | ||
|
|
f89d86944f | ||
|
|
8741d204a5 | ||
|
|
cdc85f58a8 | ||
|
|
0262d2f089 | ||
|
|
62c0343465 | ||
|
|
1e1a023fb0 | ||
|
|
1d2517ad8e | ||
|
|
d41186cb4a | ||
|
|
78e0c7eec9 | ||
|
|
1c41a94b62 | ||
|
|
2e66aafe20 | ||
|
|
55074bda76 | ||
|
|
de65bec2b7 | ||
|
|
7664dd0de3 | ||
|
|
019a88ced4 | ||
|
|
72de11abcc | ||
|
|
d71a4ebffc | ||
|
|
1089ab43bf | ||
|
|
97d4b984c9 | ||
|
|
2a8953d74d | ||
|
|
8801b10da7 | ||
|
|
6b413f2ec4 | ||
|
|
28b72694aa | ||
|
|
4afb0cfe4f | ||
|
|
3eec1281cf | ||
|
|
0660489e38 | ||
|
|
dd871a17bf | ||
|
|
dc11529862 | ||
|
|
ffabf85e31 | ||
|
|
c0026ca5ba | ||
|
|
0f2bbe71ac | ||
|
|
2a46902ecb |
@@ -0,0 +1,148 @@
|
||||
env:
|
||||
IMAGE_VERSION: "py3.12-latest"
|
||||
|
||||
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"
|
||||
watch:
|
||||
- path:
|
||||
- "fastvideo/v1/models/encoders/**"
|
||||
- "fastvideo/v1/models/loader/**"
|
||||
- "fastvideo/v1/tests/encoders/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .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/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .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/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Transformer Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=transformer
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**/*.py"
|
||||
config:
|
||||
command: "timeout 60m .buildkite/scripts/pr_test.sh"
|
||||
label: "SSIM Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=ssim
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Training Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=training
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**"
|
||||
- "csrc/attn/vsa/**"
|
||||
- "csrc/attn/tk/**"
|
||||
- "csrc/attn/setup_vsa.py"
|
||||
- "csrc/attn/config_vsa.py"
|
||||
- "csrc/attn/vsa.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Training Tests VSA"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=training_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**"
|
||||
- "csrc/attn/st_attn/**"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/attn/config_sta.py"
|
||||
- "csrc/attn/st_attn.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Inference Tests STA"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=inference_sta
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/st_attn/**"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/attn/config_sta.py"
|
||||
- "csrc/attn/st_attn.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests STA"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=precision_sta
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/vsa/**"
|
||||
- "csrc/attn/tk/**"
|
||||
- "csrc/attn/setup_vsa.py"
|
||||
- "csrc/attn/config_vsa.py"
|
||||
- "csrc/attn/vsa.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests VSA"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=precision_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
Executable
+117
@@ -0,0 +1,117 @@
|
||||
#!/bin/bash
|
||||
set -uo pipefail
|
||||
|
||||
log() {
|
||||
echo "[$(date '+%Y-%m-%d %H:%M:%S')] $1"
|
||||
}
|
||||
|
||||
log "=== Starting Modal test execution ==="
|
||||
|
||||
# Change to the project directory
|
||||
cd "$(dirname "$0")/../.."
|
||||
PROJECT_ROOT=$(pwd)
|
||||
log "Project root: $PROJECT_ROOT"
|
||||
|
||||
# Install Modal if not available
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Modal not found, installing..."
|
||||
python3 -m pip install modal
|
||||
|
||||
# Verify installation
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Error: Failed to install modal. Please install it manually."
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
|
||||
log "modal version: $(python3 -m modal --version)"
|
||||
|
||||
# Set up Modal authentication using Buildkite secrets
|
||||
log "Setting up Modal authentication from Buildkite secrets..."
|
||||
MODAL_TOKEN_ID=$(buildkite-agent secret get modal_token_id)
|
||||
MODAL_TOKEN_SECRET=$(buildkite-agent secret get modal_token_secret)
|
||||
|
||||
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
|
||||
|
||||
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
|
||||
|
||||
if [ -n "$MODAL_TOKEN_ID" ] && [ -n "$MODAL_TOKEN_SECRET" ]; then
|
||||
log "Retrieved Modal credentials from Buildkite secrets"
|
||||
python3 -m modal token set --token-id "$MODAL_TOKEN_ID" --token-secret "$MODAL_TOKEN_SECRET" --profile buildkite-ci --activate --verify
|
||||
if [ $? -eq 0 ]; then
|
||||
log "Modal authentication successful"
|
||||
else
|
||||
log "Error: Failed to set Modal credentials"
|
||||
exit 1
|
||||
fi
|
||||
else
|
||||
log "Error: Could not retrieve Modal credentials from Buildkite secrets."
|
||||
log "Please ensure 'modal_token_id' and 'modal_token_secret' secrets are set in Buildkite."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
MODAL_TEST_FILE="fastvideo/v1/tests/modal/pr_test.py"
|
||||
|
||||
if [ -z "${TEST_TYPE:-}" ]; then
|
||||
log "Error: TEST_TYPE environment variable is not set"
|
||||
exit 1
|
||||
fi
|
||||
log "Test type: $TEST_TYPE"
|
||||
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT IMAGE_VERSION=$IMAGE_VERSION"
|
||||
|
||||
case "$TEST_TYPE" in
|
||||
"encoder")
|
||||
log "Running encoder tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
|
||||
;;
|
||||
"vae")
|
||||
log "Running VAE tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
|
||||
;;
|
||||
"transformer")
|
||||
log "Running transformer tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
|
||||
;;
|
||||
"ssim")
|
||||
log "Running SSIM tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
;;
|
||||
"training")
|
||||
log "Running training tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests"
|
||||
;;
|
||||
"training_vsa")
|
||||
log "Running training VSA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests_VSA"
|
||||
;;
|
||||
"inference_sta")
|
||||
log "Running inference STA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_STA"
|
||||
;;
|
||||
"precision_sta")
|
||||
log "Running precision STA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_STA"
|
||||
;;
|
||||
"precision_vsa")
|
||||
log "Running precision VSA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_VSA"
|
||||
;;
|
||||
*)
|
||||
log "Error: Unknown test type: $TEST_TYPE"
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
|
||||
log "Executing: $MODAL_COMMAND"
|
||||
eval "$MODAL_COMMAND"
|
||||
TEST_EXIT_CODE=$?
|
||||
|
||||
if [ $TEST_EXIT_CODE -eq 0 ]; then
|
||||
log "Modal test completed successfully"
|
||||
else
|
||||
log "Error: Modal test failed with exit code: $TEST_EXIT_CODE"
|
||||
fi
|
||||
|
||||
log "=== Test execution completed with exit code: $TEST_EXIT_CODE ==="
|
||||
exit $TEST_EXIT_CODE
|
||||
@@ -0,0 +1,40 @@
|
||||
#!/bin/bash
|
||||
set -uo pipefail
|
||||
|
||||
log() {
|
||||
echo "[$(date '+%Y-%m-%d %H:%M:%S')] $1"
|
||||
}
|
||||
|
||||
log "=== Starting pre-commit checks ==="
|
||||
|
||||
cd "$(dirname "$0")/../.."
|
||||
PROJECT_ROOT=$(pwd)
|
||||
log "Project root: $PROJECT_ROOT"
|
||||
|
||||
if ! python3 -m pre_commit --version &> /dev/null; then
|
||||
log "pre-commit not found, installing..."
|
||||
python3 -m pip install --user pre-commit==4.0.1
|
||||
|
||||
if ! python3 -m pre_commit --version &> /dev/null; then
|
||||
log "Error: Failed to install pre-commit."
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
|
||||
log "Pre-commit version: $(python3 -m pre_commit --version)"
|
||||
|
||||
log "Installing/updating pre-commit hooks..."
|
||||
python3 -m pre_commit install --install-hooks
|
||||
|
||||
log "Running pre-commit checks on all files..."
|
||||
python3 -m pre_commit run --all-files
|
||||
PRE_COMMIT_EXIT_CODE=$?
|
||||
|
||||
if [ $PRE_COMMIT_EXIT_CODE -eq 0 ]; then
|
||||
log "Pre-commit checks completed successfully"
|
||||
else
|
||||
log "Error: Pre-commit checks failed with exit code: $PRE_COMMIT_EXIT_CODE"
|
||||
fi
|
||||
|
||||
log "=== Pre-commit checks completed with exit code: $PRE_COMMIT_EXIT_CODE ==="
|
||||
exit $PRE_COMMIT_EXIT_CODE
|
||||
@@ -160,8 +160,7 @@ def execute_command(pod_id):
|
||||
setup_steps = [
|
||||
"tar -xzf /tmp/repo.tar.gz --no-same-owner -C /workspace/",
|
||||
f"cd /workspace/{repo_name}",
|
||||
"source /opt/conda/etc/profile.d/conda.sh",
|
||||
"conda activate fastvideo-dev",
|
||||
"source $HOME/.local/bin/env && source /opt/venv/bin/activate",
|
||||
args.test_command
|
||||
]
|
||||
|
||||
|
||||
+209
-19
@@ -12,13 +12,11 @@ on:
|
||||
paths:
|
||||
- "fastvideo/**/*.py"
|
||||
- ".github/workflows/pr-test.yml"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
- "csrc/**"
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
custom_image:
|
||||
description: "Custom image from this repository (default: fastvideo-dev:latest)"
|
||||
required: false
|
||||
default: "fastvideo-dev:latest"
|
||||
type: string
|
||||
run_encoder_test:
|
||||
description: "Run encoder-test"
|
||||
required: false
|
||||
@@ -39,10 +37,41 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_training_test:
|
||||
description: "Run training-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_training_test_VSA:
|
||||
description: "Run training-test-VSA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_inference_test_STA:
|
||||
description: "Run inference-test-STA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_precision_test_STA:
|
||||
description: "Run precision-test-STA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_precision_test_VSA:
|
||||
description: "Run precision-test-VSA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_nightly_test:
|
||||
description: "Run nightly-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
|
||||
env:
|
||||
PYTHONUNBUFFERED: "1"
|
||||
|
||||
|
||||
concurrency:
|
||||
group: pr-test-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
@@ -59,26 +88,72 @@ jobs:
|
||||
encoder-test: ${{ steps.filter.outputs.encoder-test }}
|
||||
vae-test: ${{ steps.filter.outputs.vae-test }}
|
||||
transformer-test: ${{ steps.filter.outputs.transformer-test }}
|
||||
training-test: ${{ steps.filter.outputs.training-test }}
|
||||
training-test-VSA: ${{ steps.filter.outputs.training-test-VSA }}
|
||||
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
|
||||
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
|
||||
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
id: filter
|
||||
with:
|
||||
filters: |
|
||||
# Define reusable path patterns
|
||||
common-paths: &common-paths
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
sta-kernel-paths: &sta-kernel-paths
|
||||
- 'csrc/attn/st_attn/**'
|
||||
- 'csrc/attn/setup_sta.py'
|
||||
- 'csrc/attn/config_sta.py'
|
||||
- 'csrc/attn/st_attn.cpp'
|
||||
vsa-kernel-paths: &vsa-kernel-paths
|
||||
- 'csrc/attn/vsa/**'
|
||||
- 'csrc/attn/tk/**'
|
||||
- 'csrc/attn/setup_vsa.py'
|
||||
- 'csrc/attn/config_vsa.py'
|
||||
- 'csrc/attn/vsa.cpp'
|
||||
vsa-paths: &vsa-paths
|
||||
- 'fastvideo/v1/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
|
||||
# Actual tests
|
||||
encoder-test:
|
||||
- 'fastvideo/v1/models/encoders/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/encoders/**'
|
||||
- *common-paths
|
||||
vae-test:
|
||||
- 'fastvideo/v1/models/vaes/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/vaes/**'
|
||||
- *common-paths
|
||||
transformer-test:
|
||||
- 'fastvideo/v1/models/dits/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/transformers/**'
|
||||
- 'fastvideo/v1/layers/**'
|
||||
- 'fastvideo/v1/attention/**'
|
||||
- *common-paths
|
||||
training-test:
|
||||
- 'fastvideo/v1/**'
|
||||
- *common-paths
|
||||
training-test-VSA:
|
||||
- 'fastvideo/v1/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
inference-test-STA:
|
||||
- 'fastvideo/v1/**'
|
||||
- *common-paths
|
||||
- *sta-kernel-paths
|
||||
precision-test-STA:
|
||||
- *common-paths
|
||||
- *sta-kernel-paths
|
||||
precision-test-VSA:
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
@@ -91,8 +166,8 @@ jobs:
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -109,8 +184,8 @@ jobs:
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -127,8 +202,8 @@ jobs:
|
||||
gpu_type: "NVIDIA L40S"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -137,8 +212,7 @@ jobs:
|
||||
ssim-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
|
||||
github.event_name != 'workflow_dispatch' || (github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
@@ -155,14 +229,130 @@ jobs:
|
||||
volume_size: 200
|
||||
disk_size: 200
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:${{ matrix.python-version.tag }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
|
||||
timeout_minutes: 60
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
training-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "training-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/Vanilla -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
training-test-VSA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test-VSA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test_VSA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "training-test-VSA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/VSA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
inference-test-STA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.inference-test-STA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_inference_test_STA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "inference-test-STA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/inference/STA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
precision-test-STA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-STA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_STA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "precision-test-STA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_sta.py"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
precision-test-VSA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-VSA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_VSA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "precision-test-VSA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_block_sparse.py"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
nightly-test:
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "nightly-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
runpod-cleanup:
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
|
||||
# Add other jobs to this list as you create them
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test, training-test, training-test-VSA, inference-test-STA, precision-test-STA, precision-test-VSA]
|
||||
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
@@ -179,7 +369,7 @@ jobs:
|
||||
|
||||
- name: Cleanup all RunPod instances
|
||||
env:
|
||||
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12"]'
|
||||
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12", "training-test", "training-test-VSA", "inference-test-STA", "precision-test-STA", "precision-test-VSA"]'
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
@@ -43,6 +43,8 @@ on:
|
||||
required: true
|
||||
RUNPOD_PRIVATE_KEY:
|
||||
required: true
|
||||
WANDB_API_KEY:
|
||||
required: false
|
||||
|
||||
jobs:
|
||||
run-test:
|
||||
@@ -55,7 +57,7 @@ jobs:
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Set up SSH key
|
||||
run: |
|
||||
@@ -72,6 +74,7 @@ jobs:
|
||||
JOB_ID: ${{ inputs.job_id }}
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
timeout-minutes: ${{ inputs.timeout_minutes }}
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
|
||||
@@ -5,7 +5,7 @@ on:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/sliding_tile_attention/setup.py"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
@@ -23,13 +23,13 @@ jobs:
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/sliding_tile_attention
|
||||
cd csrc/attn
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_sta.py)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
OLD_VERSION=$(git show HEAD~1:./setup_sta.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
@@ -136,19 +136,21 @@ jobs:
|
||||
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/sliding_tile_attention # Move into the correct folder
|
||||
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
cd csrc/attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_sta.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/sliding_tile_attention
|
||||
cd csrc/attn
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-version }} | cut -d. -f1,2)
|
||||
@@ -163,7 +165,7 @@ jobs:
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}
|
||||
path: csrc/sliding_tile_attention/dist/*.whl
|
||||
path: csrc/attn/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
@@ -229,17 +231,19 @@ jobs:
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/sliding_tile_attention # Move into the correct folder
|
||||
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py sdist --dist-dir=dist
|
||||
cd csrc/attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_sta.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/sliding_tile_attention/dist/
|
||||
packages-dir: csrc/attn/dist/
|
||||
|
||||
@@ -28,4 +28,4 @@ jobs:
|
||||
|
||||
- name: Run Pytest
|
||||
run: |
|
||||
pytest --ignore csrc/sliding_tile_attention/test
|
||||
pytest --ignore csrc/attn/test
|
||||
|
||||
@@ -27,7 +27,6 @@ env
|
||||
**/build/
|
||||
**.pyc
|
||||
**.txt
|
||||
csrc/attn/tk/
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
|
||||
+2
-2
@@ -1,3 +1,3 @@
|
||||
[submodule "csrc/sliding_tile_attention/tk"]
|
||||
path = csrc/sliding_tile_attention/tk
|
||||
[submodule "csrc/attn/tk"]
|
||||
path = csrc/attn/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
@@ -91,7 +91,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
|
||||
|
||||
## Distillation and Finetuning
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/training/distillation.html)
|
||||
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetuning.html)
|
||||
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html)
|
||||
|
||||
## 📑 Development Plan
|
||||
|
||||
@@ -111,7 +111,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/developer_guide/overview.html)
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview.html)
|
||||
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects:
|
||||
|
||||
+21
-6
@@ -4,9 +4,8 @@
|
||||
|
||||
|
||||
## Installation
|
||||
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
|
||||
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:
|
||||
|
||||
```bash
|
||||
sudo apt update
|
||||
sudo apt install gcc-11 g++-11
|
||||
@@ -16,17 +15,27 @@ sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
Install STA:
|
||||
|
||||
## Environment Setup
|
||||
First, set up your CUDA environment:
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
## Install Sliding Tile Attention (STA)
|
||||
```bash
|
||||
python setup_sta.py install
|
||||
```
|
||||
|
||||
## Install Video Sparse Attention (VSA)
|
||||
```bash
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
@@ -44,8 +53,14 @@ out = sliding_tile_attention(q, k, v, window_size, 0, False)
|
||||
|
||||
## Test
|
||||
```bash
|
||||
python test/test_sta.py
|
||||
python tests/test_sta.py # test STA
|
||||
python tests/test_block_sparse.py # test VSA
|
||||
```
|
||||
## Benchmark
|
||||
```bash
|
||||
python benchmarks/bench_sta.py
|
||||
```
|
||||
|
||||
|
||||
## How Does STA Work?
|
||||
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
@@ -5,6 +5,7 @@ import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
import torch
|
||||
from st_attn import sliding_tile_attention
|
||||
from triton.testing import do_bench
|
||||
|
||||
|
||||
def flops(batch, seqlen, nheads, headdim, causal, mode="fwd"):
|
||||
@@ -13,16 +14,16 @@ def flops(batch, seqlen, nheads, headdim, causal, mode="fwd"):
|
||||
return f if mode == "fwd" else (2.5 * f if mode == "bwd" else 3.5 * f)
|
||||
|
||||
|
||||
def efficiency(flop, time):
|
||||
flop = flop / 1e12
|
||||
time = time / 1e6
|
||||
return flop / time
|
||||
def compute_TFLOPS(flops, ms):
|
||||
flops = flops / 1e12
|
||||
ms = ms / 1e3
|
||||
return flops / ms
|
||||
|
||||
|
||||
def benchmark_attention(configurations):
|
||||
results = {'fwd': defaultdict(list), 'bwd': defaultdict(list)}
|
||||
|
||||
for B, H, N, D, causal in configurations:
|
||||
for B, H, N, D, causal, dit_seq_shape, window_size in configurations:
|
||||
print("=" * 60)
|
||||
print(f"Timing forward and backward pass for B={B}, H={H}, N={N}, D={D}, causal={causal}")
|
||||
|
||||
@@ -30,38 +31,31 @@ def benchmark_attention(configurations):
|
||||
k = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
|
||||
v = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
|
||||
|
||||
grad_output = torch.randn_like(q, requires_grad=False).contiguous()
|
||||
# grad_output = torch.randn_like(q, requires_grad=False).contiguous()
|
||||
# qg = torch.zeros_like(q, requires_grad=False, dtype=torch.float).contiguous()
|
||||
# kg = torch.zeros_like(k, requires_grad=False, dtype=torch.float).contiguous()
|
||||
# vg = torch.zeros_like(v, requires_grad=False, dtype=torch.float).contiguous()
|
||||
|
||||
qg = torch.zeros_like(q, requires_grad=False, dtype=torch.float).contiguous()
|
||||
kg = torch.zeros_like(k, requires_grad=False, dtype=torch.float).contiguous()
|
||||
vg = torch.zeros_like(v, requires_grad=False, dtype=torch.float).contiguous()
|
||||
|
||||
# Prepare for timing forward pass
|
||||
start_events_fwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
|
||||
end_events_fwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
|
||||
# # Warmup for forward pass
|
||||
# for _ in range(10):
|
||||
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize()
|
||||
# # Time the forward pass
|
||||
# for i in range(10):
|
||||
# start_events_fwd[i].record()
|
||||
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
|
||||
# end_events_fwd[i].record()
|
||||
ms = do_bench(lambda: sliding_tile_attention(q, k, v, [window_size] * 24, 0, False, dit_seq_shape))
|
||||
|
||||
# Warmup for forward pass
|
||||
for _ in range(10):
|
||||
o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, '18x48x80')
|
||||
# times_fwd = [s.elapsed_time(e) for s, e in zip(start_events_fwd, end_events_fwd)]
|
||||
# time_us_fwd = np.mean(times_fwd) * 1000
|
||||
|
||||
# Time the forward pass
|
||||
for i in range(10):
|
||||
start_events_fwd[i].record()
|
||||
o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, '18x48x80')
|
||||
end_events_fwd[i].record()
|
||||
|
||||
torch.cuda.synchronize()
|
||||
times_fwd = [s.elapsed_time(e) for s, e in zip(start_events_fwd, end_events_fwd)]
|
||||
time_us_fwd = np.mean(times_fwd) * 1000
|
||||
|
||||
tflops_fwd = efficiency(flops(B, N, H, D, causal, 'fwd'), time_us_fwd)
|
||||
tflops_fwd = compute_TFLOPS(flops(B, N, H, D, causal, 'fwd'), ms)
|
||||
results['fwd'][(D, causal)].append((N, tflops_fwd))
|
||||
|
||||
print(f"Average time for forward pass in us: {time_us_fwd:.2f}")
|
||||
print(f"Average efficiency for forward pass in TFLOPS: {tflops_fwd}")
|
||||
print(f"Average time for forward pass (ms): {ms:.2f}")
|
||||
print(f"Average TFLOPS: {tflops_fwd}")
|
||||
print("-" * 60)
|
||||
|
||||
# torch.cuda.empty_cache()
|
||||
@@ -85,15 +79,14 @@ def benchmark_attention(configurations):
|
||||
# times_bwd = [s.elapsed_time(e) for s, e in zip(start_events_bwd, end_events_bwd)]
|
||||
# time_us_bwd = np.mean(times_bwd) * 1000
|
||||
|
||||
# tflops_bwd = efficiency(flops(B, N, H, D, causal, 'bwd'), time_us_bwd)
|
||||
# tflops_bwd = compute_TFLOPS(flops(B, N, H, D, causal, 'bwd'), ms)
|
||||
# results['bwd'][(D, causal)].append((N, tflops_bwd))
|
||||
|
||||
# print(f"Average time for backward pass in us: {time_us_bwd:.2f}")
|
||||
# print(f"Average efficiency for backward pass in TFLOPS: {tflops_bwd}")
|
||||
print("=" * 60)
|
||||
# print(f"Average time for backward pass(ms): {ms:.2f}")
|
||||
# print(f"Average TFLOPS: {tflops_bwd}")
|
||||
# print("=" * 60)
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
torch.cuda.synchronize()
|
||||
|
||||
return results
|
||||
|
||||
@@ -124,7 +117,10 @@ def plot_results(results):
|
||||
|
||||
# Example list of configurations to test
|
||||
configurations = [
|
||||
(2, 24, 69120, 128, False),
|
||||
(2, 24, 69120, 128, False, '18x48x80', [3, 6, 10]),
|
||||
(2, 24, 69120, 128, True, '18x48x80', [3, 6, 10]),
|
||||
(2, 24, 82944, 128, False, '36x48x48', [3, 3, 6]), # Stepvideo
|
||||
(2, 24, 82944, 128, True, '36x48x48', [3, 3, 6]),
|
||||
# (16, 16, 768*16, 128, False),
|
||||
# (16, 16, 768*2, 128, False),
|
||||
# (16, 16, 768*4, 128, False),
|
||||
@@ -4,9 +4,17 @@
|
||||
#include <cooperative_groups.h>
|
||||
#include <iostream>
|
||||
#include <stdio.h>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
// #define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
|
||||
__device__ __forceinline__ int clamp_int(int value, int min, int max) {
|
||||
return (value < min) ? min : ((value > max) ? max : value);
|
||||
}
|
||||
// #define ABS(x) ((x) < 0 ? -(x) : (x))
|
||||
__device__ __forceinline__ int abs_int(int value) {
|
||||
return (value < 0) ? -value : value;
|
||||
}
|
||||
|
||||
#define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
|
||||
#define ABS(x) ((x) < 0 ? -(x) : (x))
|
||||
|
||||
constexpr int CONSUMER_WARPGROUPS = (3);
|
||||
constexpr int PRODUCER_WARPGROUPS = (1);
|
||||
@@ -117,16 +125,16 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
|
||||
int qt = seq_idx / 6 / (CH * CW);
|
||||
int qh = (seq_idx / 6) % (CH * CW) / CW;
|
||||
int qw = (seq_idx / 6) % CW;
|
||||
qt = CLAMP(qt, DT, CT-DT-1);
|
||||
qh = CLAMP(qh, DH, CH-DH-1);
|
||||
qw = CLAMP(qw, DW, CW-DW-1);
|
||||
qt = clamp_int(qt, DT, CT-DT-1);
|
||||
qh = clamp_int(qh, DH, CH-DH-1);
|
||||
qw = clamp_int(qw, DW, CW-DW-1);
|
||||
int count = 0;
|
||||
int j = 0;
|
||||
while (count < K::stages - 1) {
|
||||
int kt = j / 3 / (CH * CW);
|
||||
int kh = (j / 3) % (CH * CW) / CW;
|
||||
int kw = (j / 3) % CW;
|
||||
bool mask = (ABS(qt - kt) <= DT) && (ABS(qh - kh) <= DH) && (ABS(qw - kw) <= DW);
|
||||
bool mask = (abs_int(qt - kt) <= DT) && (abs_int(qh - kh) <= DH) && (abs_int(qw - kw) <= DW);
|
||||
if (mask){
|
||||
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
|
||||
tma::expect_bytes(k_smem_arrived[count], sizeof(k_tile));
|
||||
@@ -167,15 +175,15 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
|
||||
int qt = seq_idx / 6 / (CH * CW);
|
||||
int qh = (seq_idx / 6) % (CH * CW) / CW;
|
||||
int qw = (seq_idx / 6) % CW;
|
||||
qt = CLAMP(qt, DT, CT-DT-1);
|
||||
qh = CLAMP(qh, DH, CH-DH-1);
|
||||
qw = CLAMP(qw, DW, CW-DW-1);
|
||||
int k_t_min = CLAMP(qt-DT, 0, CT-1);
|
||||
int k_t_max = CLAMP(qt+DT, 0, CT-1);
|
||||
int k_h_min = CLAMP(qh-DH, 0, CH-1);
|
||||
int k_h_max = CLAMP(qh+DH, 0, CH-1);
|
||||
int k_w_min = CLAMP(qw-DW, 0, CW-1);
|
||||
int k_w_max = CLAMP(qw+DW, 0, CW-1);
|
||||
qt = clamp_int(qt, DT, CT-DT-1);
|
||||
qh = clamp_int(qh, DH, CH-DH-1);
|
||||
qw = clamp_int(qw, DW, CW-DW-1);
|
||||
int k_t_min = clamp_int(qt-DT, 0, CT-1);
|
||||
int k_t_max = clamp_int(qt+DT, 0, CT-1);
|
||||
int k_h_min = clamp_int(qh-DH, 0, CH-1);
|
||||
int k_h_max = clamp_int(qh+DH, 0, CH-1);
|
||||
int k_w_min = clamp_int(qw-DW, 0, CW-1);
|
||||
int k_w_max = clamp_int(qw+DW, 0, CW-1);
|
||||
int count = 0;
|
||||
for (int kt = k_t_min; kt <= k_t_max; kt++) {
|
||||
for (int kh = k_h_min; kh <= k_h_max; kh++) {
|
||||
@@ -234,7 +242,7 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
|
||||
// the last three kv blocks are for text, we process them separately
|
||||
kv_iters = img_kv_blocks - 1;
|
||||
} else {
|
||||
kv_iters = CLAMP(DT*2+1, 1, CT) * CLAMP(DH*2+1, 1, CH) * CLAMP(DW*2+1, 1, CW) * 3 - 1 ;
|
||||
kv_iters = clamp_int(DT*2+1, 1, CT) * clamp_int(DH*2+1, 1, CH) * clamp_int(DW*2+1, 1, CW) * 3 - 1 ;
|
||||
}
|
||||
|
||||
kittens::wait(qsmem_semaphore, 0);
|
||||
@@ -415,8 +423,9 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
|
||||
float* d_l = reinterpret_cast<float*>(l_ptr);
|
||||
|
||||
cudaDeviceSynchronize();
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
//cudadevicesynchronize();
|
||||
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
|
||||
if (head_dim == 128) {
|
||||
@@ -442,8 +451,8 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
|
||||
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(text_length), static_cast<int>(hr)};
|
||||
|
||||
auto mem_size = kittens::MAX_SHARED_MEMORY;
|
||||
auto threads = NUM_WORKERS * kittens::WARP_THREADS;
|
||||
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
|
||||
int threads = NUM_WORKERS * kittens::WARP_THREADS;
|
||||
if (has_text) {
|
||||
// TORCH_CHECK(seq_len % (CONSUMER_WARPGROUPS*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 192");
|
||||
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4)-2, qo_heads, batch);
|
||||
@@ -823,10 +832,10 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
|
||||
}
|
||||
CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
cudaStreamSynchronize(stream);
|
||||
// cudaStreamSynchronize(stream);
|
||||
}
|
||||
|
||||
return o;
|
||||
cudaDeviceSynchronize();
|
||||
//cudadevicesynchronize();
|
||||
}
|
||||
|
||||
|
||||
@@ -2,11 +2,12 @@ import torch
|
||||
import argparse
|
||||
from flash_attn.utils.benchmark import benchmark_forward
|
||||
from flash_attn import flash_attn_func
|
||||
from st_attn import block_sparse_attention_fwd, block_sparse_attention_backward, BlockSparseAttentionFunction
|
||||
from st_attn import BLOCK_M, BLOCK_N
|
||||
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
|
||||
@@ -20,15 +21,6 @@ def set_seed(seed: int = 42):
|
||||
torch.cuda.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed) # if using multi-GPU
|
||||
|
||||
def parse_arguments():
|
||||
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
|
||||
parser.add_argument('--batch_size', type=int, default=4, help='Batch size')
|
||||
parser.add_argument('--num_heads', type=int, default=6, help='Number of heads')
|
||||
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=64, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[29120], help='Sequence lengths to benchmark')
|
||||
parser.add_argument('--num_iterations', type=int, default=100, help='Number of test iterations to run')
|
||||
return parser.parse_args()
|
||||
|
||||
@torch.no_grad
|
||||
def precision_metric(quant_o, fa2_o):
|
||||
@@ -135,9 +127,7 @@ def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device=
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
|
||||
|
||||
def main():
|
||||
args = parse_arguments()
|
||||
|
||||
def main(args):
|
||||
set_seed(42)
|
||||
|
||||
# Extract parameters
|
||||
@@ -191,23 +181,36 @@ def main():
|
||||
block_mask_expanded = block_sparse_mask.unsqueeze(-1).unsqueeze(-2) # [b, h, num_q_blocks, num_kv_blocks, 1, 1]
|
||||
block_mask_expanded = block_mask_expanded.expand(-1, -1, -1, -1, BLOCK_M, BLOCK_N) # [b, h, num_q_blocks, num_kv_blocks, BLOCK_M, BLOCK_N]
|
||||
full_mask = block_mask_expanded.permute(0, 1, 2, 4, 3, 5).reshape(batch, head, seq_len, seq_len)
|
||||
|
||||
q_sdpa = q.clone()
|
||||
k_sdpa = k.clone()
|
||||
v_sdpa = v.clone()
|
||||
|
||||
q.requires_grad = True
|
||||
k.requires_grad = True
|
||||
v.requires_grad = True
|
||||
q_sdpa.requires_grad = True
|
||||
k_sdpa.requires_grad = True
|
||||
v_sdpa.requires_grad = True
|
||||
|
||||
|
||||
# testing forward
|
||||
o = BlockSparseAttentionFunction.apply(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
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)
|
||||
@@ -215,52 +218,72 @@ def main():
|
||||
print(f"block_sparse_attention_fwd vs torch.nn.functional.scaled_dot_product_attention:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
# test backward
|
||||
grad_o = torch.randn_like(o)
|
||||
o.backward(grad_o)
|
||||
o_sdpa.backward(grad_o)
|
||||
|
||||
sim, l1, rmse = precision_metric(q.grad, q_sdpa.grad)
|
||||
# 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}")
|
||||
print(f"L1 error: mean={np.mean(forward_metrics['l1']):.6f}, std={np.std(forward_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(forward_metrics['rmse']):.6f}, std={np.std(forward_metrics['rmse']):.6f}")
|
||||
print(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}")
|
||||
print(f"L1 error: mean={np.mean(grad_q_metrics['l1']):.6f}, std={np.std(grad_q_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_q_metrics['rmse']):.6f}, std={np.std(grad_q_metrics['rmse']):.6f}")
|
||||
print(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}")
|
||||
print(f"L1 error: mean={np.mean(grad_k_metrics['l1']):.6f}, std={np.std(grad_k_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_k_metrics['rmse']):.6f}, std={np.std(grad_k_metrics['rmse']):.6f}")
|
||||
print(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}")
|
||||
print(f"L1 error: mean={np.mean(grad_v_metrics['l1']):.6f}, std={np.std(grad_v_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_v_metrics['rmse']):.6f}, std={np.std(grad_v_metrics['rmse']):.6f}")
|
||||
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__":
|
||||
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)
|
||||
@@ -9,14 +9,14 @@ flex_attention = torch.compile(flex_attention, dynamic=False)
|
||||
|
||||
|
||||
def flex_test(Q, K, V, kernel_size):
|
||||
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (36, 48, 48), 39, 'cuda', 0)
|
||||
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (18, 48, 80), 0, 'cuda', 0)
|
||||
output = flex_attention(Q, K, V, block_mask=mask)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def h100_fwd_kernel_test(Q, K, V, kernel_size):
|
||||
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 39, False)
|
||||
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
|
||||
return o
|
||||
|
||||
|
||||
@@ -37,7 +37,7 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
|
||||
'max_diff': 0
|
||||
},
|
||||
}
|
||||
kernel_size_ls = [(6, 1, 6), (6, 6, 1)]
|
||||
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
|
||||
from tqdm import tqdm
|
||||
for kernel_size in tqdm(kernel_size_ls):
|
||||
for _ in range(num_iterations):
|
||||
@@ -74,12 +74,14 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
|
||||
|
||||
# Example usage
|
||||
b, h, d = 2, 24, 128
|
||||
n = 82944 # Sequence length
|
||||
n = 69120 # Sequence length
|
||||
causal = False
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
# Run correctness check directly
|
||||
results = check_correctness(b, h, n, d, causal, mean, std, error_mode='output')
|
||||
assert results['TK vs FLEX']['avg_diff'] < 3e-6, f"Average difference: {results['TK vs FLEX']['avg_diff']} is too large"
|
||||
assert results['TK vs FLEX']['max_diff'] < 4e-2, f"Maximum difference: {results['TK vs FLEX']['max_diff']} is too large"
|
||||
print(f"Average difference: {results['TK vs FLEX']['avg_diff']}")
|
||||
print(f"Maximum difference: {results['TK vs FLEX']['max_diff']}")
|
||||
Submodule
+1
Submodule csrc/attn/tk added at 1719fb7264
+284
-24
@@ -2,7 +2,7 @@ import math
|
||||
|
||||
import torch
|
||||
from torch.utils.checkpoint import detach_variable
|
||||
|
||||
from typing import Tuple
|
||||
try:
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
except ImportError:
|
||||
@@ -12,33 +12,116 @@ except ImportError:
|
||||
|
||||
BLOCK_M = 64
|
||||
BLOCK_N = 64
|
||||
def block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num):
|
||||
|
||||
def video_sparse_attn(q, k, v, topk, block_size, compress_attn_weight=None):
|
||||
"""
|
||||
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.
|
||||
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]]
|
||||
"""
|
||||
# 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_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
|
||||
if isinstance(block_size, int):
|
||||
block_size = (block_size, block_size, block_size)
|
||||
|
||||
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
|
||||
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)
|
||||
|
||||
@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
|
||||
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
|
||||
v_compress)
|
||||
|
||||
output_compress = output_compress.view(batch_size, num_heads,
|
||||
seq_len // block_elements, 1,
|
||||
head_dim)
|
||||
output_compress = output_compress.repeat(1, 1, 1, block_elements,
|
||||
1).view(batch_size, num_heads,
|
||||
seq_len, head_dim)
|
||||
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num = generate_topk_block_sparse_pattern(
|
||||
block_attn_score, topk)
|
||||
|
||||
output_select = block_sparse_attn(q, k, v, q2k_block_sparse_index,
|
||||
q2k_block_sparse_num,
|
||||
k2q_block_sparse_index,
|
||||
k2q_block_sparse_num)
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
final_output = output_compress * compress_attn_weight + output_select
|
||||
else:
|
||||
final_output = output_compress + output_select
|
||||
return final_output
|
||||
|
||||
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
QK = torch.matmul(q, k.transpose(-2, -1))
|
||||
QK /= (q.size(-1)**0.5)
|
||||
|
||||
# Causal mask removed since causal is always false
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v)
|
||||
return output, QK
|
||||
|
||||
def generate_topk_block_sparse_pattern(block_attn_score: torch.Tensor,
|
||||
topk: int):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly topk kv blocks,
|
||||
based on the provided attention scores.
|
||||
|
||||
Args:
|
||||
block_attn_score: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Attention scores between query and key blocks
|
||||
topk: int
|
||||
Number of kv blocks each q block attends to
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, topk]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to topk).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, max_q_per_kv]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
"""
|
||||
device = block_attn_score.device
|
||||
# Extract dimensions from block_attn_score
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_attn_score.shape
|
||||
|
||||
sorted_result = torch.sort(block_attn_score, dim=-1, descending=True)
|
||||
|
||||
sorted_indice = sorted_result.indices
|
||||
|
||||
q2k_block_sparse_index, _ = torch.sort(sorted_indice[:, :, :, :topk],
|
||||
dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(dtype=torch.int32)
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks),
|
||||
topk,
|
||||
device=device,
|
||||
dtype=torch.int32)
|
||||
|
||||
block_map = topk_index_to_map(q2k_block_sparse_index,
|
||||
num_kv_blocks,
|
||||
transpose_map=True)
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(
|
||||
block_map.transpose(2, 3))
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
|
||||
@torch._dynamo.disable
|
||||
def block_sparse_attn(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
@@ -61,6 +144,20 @@ def block_sparse_attn(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q
|
||||
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
|
||||
@@ -128,6 +225,169 @@ def index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, BLOCK_Q, BLOCK_K
|
||||
|
||||
return mask
|
||||
|
||||
@triton.jit
|
||||
def topk_index_to_map_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
topk: tl.constexpr,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
for i in tl.static_range(topk):
|
||||
index = tl.load(index_ptr_base + i * index_kv_stride)
|
||||
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
|
||||
|
||||
@triton.jit
|
||||
def map_to_index_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
index_num_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
index_num_bs_stride,
|
||||
index_num_h_stride,
|
||||
index_num_q_stride,
|
||||
num_kv_blocks: tl.constexpr,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
num = 0
|
||||
for i in tl.static_range(num_kv_blocks):
|
||||
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
|
||||
if map_entry:
|
||||
tl.store(index_ptr_base + num * index_kv_stride, i)
|
||||
num += 1
|
||||
|
||||
tl.store(
|
||||
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
|
||||
q * index_num_q_stride, num)
|
||||
|
||||
def topk_index_to_map(index: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
transpose_map: bool = False):
|
||||
"""
|
||||
Convert topk indices to a map.
|
||||
|
||||
Args:
|
||||
index: [bs, h, num_q_blocks, topk]
|
||||
The topk indices tensor.
|
||||
num_kv_blocks: int
|
||||
The number of key-value blocks in the block_map returned
|
||||
transpose_map: bool
|
||||
If True, the block_map will be transposed on the final two dimensions.
|
||||
|
||||
Returns:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
A binary map where 1 indicates that the q block attends to the kv block.
|
||||
"""
|
||||
bs, h, num_q_blocks, topk = index.shape
|
||||
|
||||
if transpose_map is False:
|
||||
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
else:
|
||||
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
block_map = block_map.transpose(2, 3)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
topk_index_to_map_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
topk=topk,
|
||||
)
|
||||
|
||||
return block_map
|
||||
|
||||
def map_to_index(block_map: torch.Tensor):
|
||||
"""
|
||||
Convert a block map to indices and counts.
|
||||
|
||||
Args:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The block map tensor.
|
||||
|
||||
Returns:
|
||||
index: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The indices of the blocks.
|
||||
index_num: [bs, h, num_q_blocks]
|
||||
The number of blocks for each q block.
|
||||
"""
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
|
||||
|
||||
index = torch.full((block_map.shape),
|
||||
-1,
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
index_num = torch.empty((bs, h, num_q_blocks),
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
map_to_index_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
index_num,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
index_num.stride(0),
|
||||
index_num.stride(1),
|
||||
index_num.stride(2),
|
||||
num_kv_blocks=num_kv_blocks,
|
||||
)
|
||||
|
||||
return index, index_num
|
||||
|
||||
class BlockSparseAttentionFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
o, lse = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
ctx.save_for_backward(q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
return o
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num = ctx.saved_tensors
|
||||
grad_q, grad_k, grad_v = block_sparse_attention_backward(
|
||||
q, k, v, o, lse, grad_output, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
)
|
||||
return grad_q, grad_k, grad_v, None, None, None, None
|
||||
|
||||
|
||||
class DummyOperator(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, x):
|
||||
@@ -207,4 +467,4 @@ class BlockSparseAttnTorch:
|
||||
|
||||
o = DummyOperator.apply(output)
|
||||
o.register_hook(self.recompute)
|
||||
return o
|
||||
return o
|
||||
|
||||
@@ -3,6 +3,8 @@
|
||||
#include "kittens.cuh"
|
||||
#include <cooperative_groups.h>
|
||||
#include <iostream>
|
||||
#include <c10/cuda/CUDAGuard.h>
|
||||
|
||||
|
||||
using namespace kittens;
|
||||
namespace cg = cooperative_groups;
|
||||
@@ -940,8 +942,9 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
|
||||
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
|
||||
float* d_l = reinterpret_cast<float*>(l_ptr);
|
||||
|
||||
cudaDeviceSynchronize();
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
//cudadevicesynchronize();
|
||||
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
if (head_dim == 64) {
|
||||
using q_tile = st_bf<fwd_attend_ker_tile_dims<64>::qo_height, fwd_attend_ker_tile_dims<64>::tile_width>;
|
||||
@@ -966,7 +969,7 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
|
||||
|
||||
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())};
|
||||
|
||||
auto mem_size = 54000;
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
dim3 grid(seq_len/(64), qo_heads, batch);
|
||||
|
||||
@@ -979,7 +982,7 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
|
||||
fwd_attend_ker<64><<<grid, (128), mem_size, stream>>>(g);
|
||||
|
||||
CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
cudaStreamSynchronize(stream);
|
||||
// cudaStreamSynchronize(stream);
|
||||
}
|
||||
|
||||
if (head_dim == 128) {
|
||||
@@ -1005,7 +1008,7 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
|
||||
|
||||
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())};
|
||||
|
||||
auto mem_size = 54000;
|
||||
constexpr int mem_size = 54000;
|
||||
|
||||
dim3 grid(seq_len/(64), qo_heads, batch);
|
||||
|
||||
@@ -1018,11 +1021,11 @@ block_sparse_attention_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v
|
||||
fwd_attend_ker<128><<<grid, (128), mem_size, stream>>>(g);
|
||||
|
||||
CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
cudaStreamSynchronize(stream);
|
||||
// cudaStreamSynchronize(stream);
|
||||
}
|
||||
|
||||
return {o, l_vec};
|
||||
cudaDeviceSynchronize();
|
||||
//cudadevicesynchronize();
|
||||
}
|
||||
|
||||
std::vector<torch::Tensor>
|
||||
@@ -1132,13 +1135,14 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
float* d_kg = reinterpret_cast<float*>(kg_ptr);
|
||||
float* d_vg = reinterpret_cast<float*>(vg_ptr);
|
||||
|
||||
auto mem_size = kittens::MAX_SHARED_MEMORY;
|
||||
auto threads = 4 * kittens::WARP_THREADS;
|
||||
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
|
||||
int threads = 4 * kittens::WARP_THREADS;
|
||||
|
||||
cudaDeviceSynchronize();
|
||||
auto stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
//cudadevicesynchronize();
|
||||
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
|
||||
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
|
||||
|
||||
cudaStreamSynchronize(stream);
|
||||
// cudaStreamSynchronize(stream);
|
||||
|
||||
// TORCH_CHECK(seq_len % (4*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 256");
|
||||
dim3 grid_bwd(seq_len/(4*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
|
||||
@@ -1222,7 +1226,7 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
|
||||
threads = 128;
|
||||
|
||||
cudaDeviceSynchronize();
|
||||
//cudadevicesynchronize();
|
||||
|
||||
{
|
||||
cudaFuncSetAttribute(
|
||||
@@ -1240,8 +1244,8 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
}
|
||||
|
||||
// CHECK_CUDA_ERROR(cudaGetLastError());
|
||||
cudaStreamSynchronize(stream);
|
||||
cudaDeviceSynchronize();
|
||||
// cudaStreamSynchronize(stream);
|
||||
//cudadevicesynchronize();
|
||||
// const auto kernel_end = std::chrono::high_resolution_clock::now();
|
||||
// std::cout << "Kernel Time: " << std::chrono::duration_cast<std::chrono::microseconds>(kernel_end - start).count() << "us" << std::endl;
|
||||
// std::cout << "---" << std::endl;
|
||||
@@ -1326,7 +1330,7 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
|
||||
threads = 128;
|
||||
|
||||
cudaDeviceSynchronize();
|
||||
//cudadevicesynchronize();
|
||||
|
||||
{
|
||||
cudaFuncSetAttribute(
|
||||
@@ -1338,10 +1342,10 @@ block_sparse_attention_backward(torch::Tensor q,
|
||||
bwd_attend_ker<128><<<grid_bwd_2, threads, 113000, stream>>>(bwd_global);
|
||||
}
|
||||
|
||||
cudaStreamSynchronize(stream);
|
||||
cudaDeviceSynchronize();
|
||||
// cudaStreamSynchronize(stream);
|
||||
//cudadevicesynchronize();
|
||||
}
|
||||
|
||||
return {qg, kg, vg};
|
||||
cudaDeviceSynchronize();
|
||||
//cudadevicesynchronize();
|
||||
}
|
||||
@@ -1,7 +1,9 @@
|
||||
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
|
||||
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
@@ -9,17 +11,25 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
zsh \
|
||||
vim \
|
||||
curl \
|
||||
gcc-11 \
|
||||
g++-11 \
|
||||
clang-11 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
|
||||
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
|
||||
rm Miniconda3-latest-Linux-x86_64.sh
|
||||
# Set up C++20 compilers for ThunderKittens
|
||||
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
ENV PATH=/opt/conda/bin:$PATH
|
||||
# Set CUDA environment variables
|
||||
ENV CUDA_HOME=/usr/local/cuda-12.8
|
||||
ENV PATH=${CUDA_HOME}/bin:${PATH}
|
||||
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
RUN conda create --name fastvideo-dev python=3.12.9 -y
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
# Install uv and source its environment
|
||||
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
|
||||
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject.toml ./
|
||||
@@ -27,22 +37,36 @@ COPY pyproject.toml ./
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
|
||||
conda clean -afy
|
||||
# Create and activate virtual environment with specific Python version and seed
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.12 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir flash-attn==2.8.0.post2 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Remove authentication headers
|
||||
RUN git config --unset-all http.https://github.com/.extraheader || true
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_sta.py install
|
||||
|
||||
# Set up automatic conda environment activation for all shells
|
||||
RUN echo 'source /opt/conda/etc/profile.d/conda.sh' >> /root/.bashrc && \
|
||||
echo 'conda activate fastvideo-dev' >> /root/.bashrc && \
|
||||
# Ensure .bashrc is sourced for SSH login shells
|
||||
echo 'if [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_vsa.py install
|
||||
|
||||
EXPOSE 22
|
||||
@@ -1,7 +1,7 @@
|
||||
(sta-demo)=
|
||||
|
||||
# 🔍 Demo
|
||||
There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
This is is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
<div style="text-align: center;">
|
||||
<video controls width="800">
|
||||
@@ -9,3 +9,9 @@ There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
Your browser does not support the video tag.
|
||||
</video>
|
||||
</div>
|
||||
|
||||
You can run STA using the following command:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_STA.sh
|
||||
```
|
||||
|
||||
@@ -7,70 +7,40 @@ To save GPU memory, we precompute text embeddings and VAE latents to eliminate t
|
||||
We provide a sample dataset to help you get started. Download the source media using the following command:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Image-Vid-Finetune-Src --local_dir=data/Image-Vid-Finetune-Src --repo_type=dataset
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/mini_i2v_dataset --local_dir=FastVideo/mini_i2v_dataset --repo_type=dataset
|
||||
```
|
||||
|
||||
The folder `crush-smol_raw/` contains raw videos and captions for testing preprocessing, while `crush-smol_preprocessed/` contains latents prepared for testing training.
|
||||
|
||||
To preprocess the dataset for fine-tuning or distillation, run:
|
||||
|
||||
```
|
||||
bash scripts/preprocess/preprocess_mochi_data.sh # for mochi
|
||||
bash scripts/preprocess/preprocess_hunyuan_data.sh # for hunyuan
|
||||
bash scripts/preprocess/v1_preprocess_wan_data_t2v # for wan
|
||||
```
|
||||
|
||||
The preprocessed dataset will be stored in `Image-Vid-Finetune-Mochi` or `Image-Vid-Finetune-HunYuan` correspondingly.
|
||||
|
||||
## Process your own dataset
|
||||
|
||||
If you wish to create your own dataset for finetuning or distillation, please structure you video dataset in the following format:
|
||||
If you wish to create your own dataset for finetuning or distillation, please refer `mini_i2v_dataset/crush-smol_raw/` to structure you video dataset in the following format:
|
||||
|
||||
```
|
||||
path_to_dataset_folder/
|
||||
├── media/
|
||||
│ ├── 0.jpg
|
||||
path_to_your_dataset_folder/
|
||||
├── videos/
|
||||
│ ├── 0.mp4
|
||||
│ ├── 1.mp4
|
||||
│ ├── 2.jpg
|
||||
├── video2caption.json
|
||||
└── merge.txt
|
||||
├── videos.txt
|
||||
└── prompt.txt
|
||||
```
|
||||
|
||||
Format the JSON file as a list, where each item represents a media source:
|
||||
To geranate the `videos2caption.json` and `merge.txt`, run
|
||||
|
||||
For image media,
|
||||
|
||||
```
|
||||
{
|
||||
"path": "0.jpg",
|
||||
"cap": ["captions"]
|
||||
}
|
||||
``` python
|
||||
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
|
||||
```
|
||||
|
||||
For video media,
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/v1_preprocess_****.sh` accordingly and run:
|
||||
|
||||
```
|
||||
{
|
||||
"path": "1.mp4",
|
||||
"resolution": {
|
||||
"width": 848,
|
||||
"height": 480
|
||||
},
|
||||
"fps": 30.0,
|
||||
"duration": 6.033333333333333,
|
||||
"cap": [
|
||||
"caption"
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
Use a txt file (merge.txt) to contain the source folder for media and the JSON file for meta information:
|
||||
|
||||
```
|
||||
path_to_media_source_foder,path_to_json_file
|
||||
```
|
||||
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/preprocess_****_data.sh` accordingly and run:
|
||||
|
||||
```
|
||||
bash scripts/preprocess/preprocess_****_data.sh
|
||||
bash scripts/preprocess/v1_preprocess_****.sh
|
||||
```
|
||||
|
||||
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
|
||||
|
||||
@@ -16,6 +16,13 @@ bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
## ⚡ Finetune with VSA
|
||||
Follow [data_preprocess.md](#v0-data-preprocess) to get parquet files for preproccessed latent, and then run:
|
||||
|
||||
```bash
|
||||
bash scripts/finetune/finetune_v1_VSA.sh
|
||||
```
|
||||
|
||||
## ⚡ Lora Finetune
|
||||
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
|
||||
@@ -10,7 +10,7 @@ def main():
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
# if num_gpus > 1, FastVideo will automatically handle distributed setup
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=True,
|
||||
use_cpu_offload=False
|
||||
|
||||
@@ -5,7 +5,7 @@ export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
|
||||
export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
|
||||
|
||||
base_port=29503
|
||||
num_gpu=$(nvidia-smi --query-gpu=gpu_name --format=csv,noheader | wc -l)
|
||||
num_gpu=1
|
||||
gpu_ids=$(seq 0 $((num_gpu-1)))
|
||||
skip_time_steps=12
|
||||
|
||||
@@ -14,7 +14,7 @@ STA_mode="STA_searching"
|
||||
for i in $gpu_ids; do
|
||||
port=$((base_port+i))
|
||||
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
|
||||
--prompt_path ./assets/prompt_extend_${i}.txt \
|
||||
--prompt_path ./assets/prompt_${i}.txt \
|
||||
--output_path $output_path \
|
||||
--STA_mode $STA_mode &
|
||||
sleep 1
|
||||
@@ -27,7 +27,7 @@ STA_mode="STA_tuning"
|
||||
for i in $gpu_ids; do
|
||||
port=$((base_port+i))
|
||||
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
|
||||
--prompt_path ./assets/prompt_extend_${i}.txt \
|
||||
--prompt_path ./assets/prompt_${i}.txt \
|
||||
--output_path $output_path \
|
||||
--STA_mode $STA_mode \
|
||||
--skip_time_steps $skip_time_steps &
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
This directory contain e2e examples scripts for finetuning Wan2.1 I2V.
|
||||
|
||||
Execute the following commands from `FastVideo/` to run training:
|
||||
|
||||
- Download crush-smol dataset:
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/download_dataset.sh`
|
||||
- Preprocess the videos and captions into latents:
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/preprocess_wan_data_i2v.sh`
|
||||
- Edit the following file and run finetuning:
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/finetune_i2v.sh`
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,91 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_i2v/combined_parquet_dataset/"
|
||||
VALIDATION_DIR="data/crush-smol_processed_i2v/validation_parquet_dataset/"
|
||||
NUM_GPUS=8
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_i2v_finetune"
|
||||
--output_dir "$DATA_DIR/outputs/wan_i2v_finetune"
|
||||
--max_train_steps 2000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 8
|
||||
--tp_size 8
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_preprocessed_path "$VALIDATION_DIR"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,130 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=i2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --qos=hao
|
||||
#SBATCH --nodes=4
|
||||
#SBATCH --ntasks=4
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=i2v_output/i2v_%j.out
|
||||
#SBATCH --error=i2v_output/i2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_i2v/combined_parquet_dataset/"
|
||||
VALIDATION_DIR="data/crush-smol_processed_i2v/validation_parquet_dataset/"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_i2v_finetune
|
||||
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune_2n"
|
||||
--max_train_steps=2000
|
||||
--train_batch_size=2
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps=1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size $NUM_GPUS
|
||||
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 10
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_preprocessed_path "$VALIDATION_DIR"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate=1e-5
|
||||
--mixed_precision="bf16"
|
||||
--checkpointing_steps=1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,25 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_i2v/"
|
||||
VALIDATION_PATH="examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation.json"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 77 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--model_type $MODEL_TYPE \
|
||||
--train_fps 16 \
|
||||
--validation_dataset_file $VALIDATION_PATH \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--preprocess_task "i2v"
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -0,0 +1,10 @@
|
||||
This directory contain e2e examples scripts for finetuning Wan2.1 T2v.
|
||||
|
||||
Execute the following commands from `FastVideo/` to run training:
|
||||
|
||||
- Download crush-smol dataset:
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/download_dataset.sh`
|
||||
- Preprocess the videos and captions into latents:
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/preprocess_wan_data_t2v.sh`
|
||||
- Edit the following file and run finetuning:
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/finetune_t2v.sh`
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,90 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
|
||||
VALIDATION_DIR="data/crush-smol_processed_t2v/validation_parquet_dataset/"
|
||||
NUM_GPUS=4
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_t2v_finetune"
|
||||
--output_dir "outputs/wan_t2v_finetune"
|
||||
--max_train_steps 5000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 8
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size $NUM_GPUS
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path $DATA_DIR
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_preprocessed_path $VALIDATION_DIR
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 6000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,127 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --qos=hao
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=t2v_output/t2v_%j.out
|
||||
#SBATCH --error=t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
|
||||
VALIDATION_DIR="data/crush-smol_processed_t2v/validation_parquet_dataset/"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_finetune
|
||||
--output_dir="outputs/wan_t2v_finetune"
|
||||
--max_train_steps=1000
|
||||
--train_batch_size=4
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps=1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 4
|
||||
--tp_size 4
|
||||
--hsdp_replicate_dim 2
|
||||
--hsdp_shard_dim 4
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 10
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_preprocessed_path "$VALIDATION_DIR"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate=5e-5
|
||||
--mixed_precision="bf16"
|
||||
--checkpointing_steps=500
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/v1/training/wan_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,25 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v/"
|
||||
VALIDATION_PATH="examples/training/finetune/wan_t2v_1_3b/crush_smol/validation.json"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 77 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--model_type $MODEL_TYPE \
|
||||
--train_fps 16 \
|
||||
--validation_dataset_file $VALIDATION_PATH \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--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": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -104,13 +104,7 @@ if __name__ == "__main__":
|
||||
default=None,
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
|
||||
@@ -671,13 +671,6 @@ if __name__ == "__main__":
|
||||
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
parser.add_argument("--num_train_epochs", type=int, default=100)
|
||||
|
||||
@@ -693,13 +693,6 @@ if __name__ == "__main__":
|
||||
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
parser.add_argument("--num_train_epochs", type=int, default=100)
|
||||
|
||||
+1
-7
@@ -520,13 +520,7 @@ if __name__ == "__main__":
|
||||
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
parser.add_argument("--num_train_epochs", type=int, default=100)
|
||||
|
||||
@@ -6,6 +6,8 @@ from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.v1.utils import dict_to_3d_list
|
||||
|
||||
|
||||
def configure_sta(mode: str = 'STA_searching',
|
||||
layer_num: int = 40,
|
||||
@@ -349,21 +351,6 @@ def select_best_mask_strategy(
|
||||
return best_mask_strategy, overall_sparsity, strategy_counts
|
||||
|
||||
|
||||
def dict_to_3d_list(mask_strategy: Optional[Dict[str, List[int]]],
|
||||
t_max: int = 50,
|
||||
l_max: int = 60,
|
||||
h_max: int = 24) -> List[List[List[Optional[List[int]]]]]:
|
||||
result: List[List[List[Optional[List[int]]]]] = [[[
|
||||
None for _ in range(h_max)
|
||||
] for _ in range(l_max)] for _ in range(t_max)]
|
||||
if mask_strategy is None:
|
||||
return result
|
||||
for key, value in mask_strategy.items():
|
||||
t, layer_idx, h = map(int, key.split('_'))
|
||||
result[t][layer_idx][h] = value
|
||||
return result
|
||||
|
||||
|
||||
def save_mask_search_results(
|
||||
mask_search_final_result: List[Dict[str, List[float]]],
|
||||
prompt: str,
|
||||
|
||||
@@ -10,8 +10,8 @@ from fastvideo.v1.attention.selector import get_attn_backend
|
||||
|
||||
__all__ = [
|
||||
"DistributedAttention",
|
||||
"DistributedAttention_VSA",
|
||||
"LocalAttention",
|
||||
"DistributedAttention_VSA",
|
||||
"AttentionBackend",
|
||||
"AttentionMetadata",
|
||||
"AttentionMetadataBuilder",
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Type
|
||||
from typing import Any, List, Optional, Type
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
@@ -17,33 +17,11 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.utils import dict_to_3d_list
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
# TODO(will-refactor): move this to a utils file
|
||||
def dict_to_3d_list(
|
||||
mask_strategy: Dict[str,
|
||||
Any]) -> List[List[List[Optional[torch.Tensor]]]]:
|
||||
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
|
||||
|
||||
max_timesteps_idx = max(
|
||||
timesteps_idx for timesteps_idx, layer_idx, head_idx in indices) + 1
|
||||
max_layer_idx = max(layer_idx
|
||||
for timesteps_idx, layer_idx, head_idx in indices) + 1
|
||||
max_head_idx = max(head_idx
|
||||
for timesteps_idx, layer_idx, head_idx in indices) + 1
|
||||
|
||||
result = [[[None for _ in range(max_head_idx)]
|
||||
for _ in range(max_layer_idx)] for _ in range(max_timesteps_idx)]
|
||||
|
||||
for key, value in mask_strategy.items():
|
||||
timesteps_idx, layer_idx, head_idx = map(int, key.split('_'))
|
||||
result[timesteps_idx][layer_idx][head_idx] = value
|
||||
|
||||
return result
|
||||
|
||||
|
||||
class RangeDict(dict):
|
||||
|
||||
def __getitem__(self, item: int) -> str:
|
||||
@@ -287,16 +265,10 @@ class SlidingTileAttentionImpl(AttentionImpl):
|
||||
forward_batch.mask_search_final_result_pos[timestep].append(
|
||||
layer_loss_save)
|
||||
else:
|
||||
# windows = [
|
||||
# self.mask_strategy[timestep][layer_idx][head_idx + start_head]
|
||||
# for head_idx in range(head_num)
|
||||
# ]
|
||||
windows = [
|
||||
STA_param[head_idx + start_head] for head_idx in range(head_num)
|
||||
]
|
||||
# if has_text is False:
|
||||
# from IPython import embed
|
||||
# embed()
|
||||
|
||||
hidden_states = sliding_tile_attention(
|
||||
query, key, value, windows, text_length, has_text,
|
||||
self.dit_seq_shape_str).transpose(1, 2)
|
||||
|
||||
@@ -1,18 +1,15 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import math
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, List, Optional, Type, cast
|
||||
from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
import triton
|
||||
import triton.language as tl
|
||||
from einops import rearrange
|
||||
|
||||
try:
|
||||
from vsa import block_sparse_attn
|
||||
except ImportError: # noqa: E722
|
||||
block_sparse_attn = None
|
||||
from typing import Tuple
|
||||
from vsa import video_sparse_attn
|
||||
except ImportError:
|
||||
video_sparse_attn = None
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
@@ -75,14 +72,18 @@ class VideoSparseAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
if forward_batch.latents is None:
|
||||
raise ValueError("latents cannot be None")
|
||||
|
||||
raw_latent_shape = forward_batch.latents.shape
|
||||
patch_size = fastvideo_args.dit_config.patch_size
|
||||
raw_latent_shape = forward_batch.raw_latent_shape
|
||||
if raw_latent_shape is None:
|
||||
raise ValueError("raw_latent_shape cannot be None")
|
||||
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.patch_size
|
||||
dit_seq_shape = [
|
||||
raw_latent_shape[2] // patch_size[0],
|
||||
raw_latent_shape[3] // patch_size[1],
|
||||
raw_latent_shape[4] // patch_size[2]
|
||||
]
|
||||
VSA_sparsity = forward_batch.VSA_sparsity
|
||||
|
||||
return VideoSparseAttentionMetadata(current_timestep=current_timestep,
|
||||
dit_seq_shape=dit_seq_shape,
|
||||
VSA_sparsity=VSA_sparsity)
|
||||
@@ -177,12 +178,16 @@ class VideoSparseAttentionImpl(AttentionImpl):
|
||||
value = value.transpose(1, 2).contiguous()
|
||||
gate_compress = gate_compress.transpose(1, 2).contiguous()
|
||||
|
||||
VSA_sparsity = attn_metadata.VSA_sparsity
|
||||
|
||||
cur_topk = math.ceil(
|
||||
(1 - attn_metadata.VSA_sparsity) *
|
||||
(1 - VSA_sparsity) *
|
||||
(self.img_seq_length / math.prod(self.VSA_base_tile_size)))
|
||||
|
||||
# Cast to Any to bypass type checking for untyped function
|
||||
hidden_states = cast(Any, sparse_attn_c_s_p)(
|
||||
if video_sparse_attn is None:
|
||||
raise NotImplementedError("video_sparse_attn is not installed")
|
||||
|
||||
hidden_states = video_sparse_attn(
|
||||
query,
|
||||
key,
|
||||
value,
|
||||
@@ -191,267 +196,3 @@ class VideoSparseAttentionImpl(AttentionImpl):
|
||||
compress_attn_weight=gate_compress).transpose(1, 2)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
QK = torch.matmul(q, k.transpose(-2, -1))
|
||||
QK /= (q.size(-1)**0.5)
|
||||
|
||||
# Causal mask removed since causal is always false
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v)
|
||||
return output, QK
|
||||
|
||||
|
||||
def sparse_attn_c_s_p(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
|
||||
|
||||
|
||||
@triton.jit
|
||||
def topk_index_to_map_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
topk: tl.constexpr,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
for i in tl.static_range(topk):
|
||||
index = tl.load(index_ptr_base + i * index_kv_stride)
|
||||
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
|
||||
|
||||
|
||||
@triton.jit
|
||||
def map_to_index_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
index_num_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
index_num_bs_stride,
|
||||
index_num_h_stride,
|
||||
index_num_q_stride,
|
||||
num_kv_blocks: tl.constexpr,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
num = 0
|
||||
for i in tl.static_range(num_kv_blocks):
|
||||
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
|
||||
if map_entry:
|
||||
tl.store(index_ptr_base + num * index_kv_stride, i)
|
||||
num += 1
|
||||
|
||||
tl.store(
|
||||
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
|
||||
q * index_num_q_stride, num)
|
||||
|
||||
|
||||
def topk_index_to_map(index: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
transpose_map: bool = False):
|
||||
"""
|
||||
Convert topk indices to a map.
|
||||
|
||||
Args:
|
||||
index: [bs, h, num_q_blocks, topk]
|
||||
The topk indices tensor.
|
||||
num_kv_blocks: int
|
||||
The number of key-value blocks in the block_map returned
|
||||
transpose_map: bool
|
||||
If True, the block_map will be transposed on the final two dimensions.
|
||||
|
||||
Returns:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
A binary map where 1 indicates that the q block attends to the kv block.
|
||||
"""
|
||||
bs, h, num_q_blocks, topk = index.shape
|
||||
|
||||
if transpose_map is False:
|
||||
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
else:
|
||||
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
block_map = block_map.transpose(2, 3)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
topk_index_to_map_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
topk=topk,
|
||||
)
|
||||
|
||||
return block_map
|
||||
|
||||
|
||||
def map_to_index(block_map: torch.Tensor):
|
||||
"""
|
||||
Convert a block map to indices and counts.
|
||||
|
||||
Args:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The block map tensor.
|
||||
|
||||
Returns:
|
||||
index: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The indices of the blocks.
|
||||
index_num: [bs, h, num_q_blocks]
|
||||
The number of blocks for each q block.
|
||||
"""
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
|
||||
|
||||
index = torch.full((block_map.shape),
|
||||
-1,
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
index_num = torch.empty((bs, h, num_q_blocks),
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
map_to_index_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
index_num,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
index_num.stride(0),
|
||||
index_num.stride(1),
|
||||
index_num.stride(2),
|
||||
num_kv_blocks=num_kv_blocks,
|
||||
)
|
||||
|
||||
return index, index_num
|
||||
|
||||
|
||||
def generate_topk_block_sparse_pattern(block_attn_score: torch.Tensor,
|
||||
topk: int):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly topk kv blocks,
|
||||
based on the provided attention scores.
|
||||
|
||||
Args:
|
||||
block_attn_score: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Attention scores between query and key blocks
|
||||
topk: int
|
||||
Number of kv blocks each q block attends to
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, topk]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to topk).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, max_q_per_kv]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
"""
|
||||
device = block_attn_score.device
|
||||
# Extract dimensions from block_attn_score
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_attn_score.shape
|
||||
|
||||
sorted_result = torch.sort(block_attn_score, dim=-1, descending=True)
|
||||
|
||||
sorted_indice = sorted_result.indices
|
||||
|
||||
q2k_block_sparse_index, _ = torch.sort(sorted_indice[:, :, :, :topk],
|
||||
dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(dtype=torch.int32)
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks),
|
||||
topk,
|
||||
device=device,
|
||||
dtype=torch.int32)
|
||||
|
||||
block_map = topk_index_to_map(q2k_block_sparse_index,
|
||||
num_kv_blocks,
|
||||
transpose_map=True)
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(
|
||||
block_map.transpose(2, 3))
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
|
||||
@@ -12,7 +12,7 @@ from fastvideo.v1.distributed.communication_op import (
|
||||
from fastvideo.v1.distributed.parallel_state import (get_sp_parallel_rank,
|
||||
get_sp_world_size)
|
||||
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
from fastvideo.v1.utils import get_compute_dtype
|
||||
|
||||
|
||||
@@ -26,8 +26,8 @@ class DistributedAttention(nn.Module):
|
||||
num_kv_heads: Optional[int] = None,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[
|
||||
AttentionBackendEnum, ...]] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
@@ -45,13 +45,13 @@ class DistributedAttention(nn.Module):
|
||||
dtype,
|
||||
supported_attention_backends=supported_attention_backends)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
causal=causal,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
prefix=f"{prefix}.impl",
|
||||
**extra_impl_args)
|
||||
self.attn_impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
causal=causal,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
prefix=f"{prefix}.impl",
|
||||
**extra_impl_args)
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
self.num_kv_heads = num_kv_heads
|
||||
@@ -100,7 +100,7 @@ class DistributedAttention(nn.Module):
|
||||
scatter_dim=2,
|
||||
gather_dim=1)
|
||||
# Apply backend-specific preprocess_qkv
|
||||
qkv = self.impl.preprocess_qkv(qkv, ctx_attn_metadata)
|
||||
qkv = self.attn_impl.preprocess_qkv(qkv, ctx_attn_metadata)
|
||||
|
||||
# Concatenate with replicated QKV if provided
|
||||
if replicated_q is not None:
|
||||
@@ -116,7 +116,7 @@ class DistributedAttention(nn.Module):
|
||||
|
||||
q, k, v = qkv.chunk(3, dim=0)
|
||||
|
||||
output = self.impl.forward(q, k, v, ctx_attn_metadata)
|
||||
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
||||
|
||||
# Redistribute back if using sequence parallelism
|
||||
replicated_output = None
|
||||
@@ -127,7 +127,7 @@ class DistributedAttention(nn.Module):
|
||||
replicated_output = sequence_model_parallel_all_gather(
|
||||
replicated_output.contiguous(), dim=2)
|
||||
# Apply backend-specific postprocess_output
|
||||
output = self.impl.postprocess_output(output, ctx_attn_metadata)
|
||||
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
|
||||
|
||||
output = sequence_model_parallel_all_to_all_4D(output,
|
||||
scatter_dim=1,
|
||||
@@ -183,18 +183,17 @@ class DistributedAttention_VSA(DistributedAttention):
|
||||
scatter_dim=2,
|
||||
gather_dim=1)
|
||||
|
||||
qkvg = self.impl.preprocess_qkv(
|
||||
qkvg, ctx_attn_metadata) # (yongqi) pass latent shape here?
|
||||
qkvg = self.attn_impl.preprocess_qkv(qkvg, ctx_attn_metadata)
|
||||
|
||||
q, k, v, gate_compress = qkvg.chunk(4, dim=0)
|
||||
output = self.impl.forward(q, k, v, gate_compress,
|
||||
ctx_attn_metadata) # type: ignore[call-arg]
|
||||
output = self.attn_impl.forward(
|
||||
q, k, v, gate_compress, ctx_attn_metadata) # type: ignore[call-arg]
|
||||
|
||||
# Redistribute back if using sequence parallelism
|
||||
replicated_output = None
|
||||
|
||||
# Apply backend-specific postprocess_output
|
||||
output = self.impl.postprocess_output(output, ctx_attn_metadata)
|
||||
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
|
||||
|
||||
output = sequence_model_parallel_all_to_all_4D(output,
|
||||
scatter_dim=1,
|
||||
@@ -212,8 +211,8 @@ class LocalAttention(nn.Module):
|
||||
num_kv_heads: Optional[int] = None,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[
|
||||
AttentionBackendEnum, ...]] = None,
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
if softmax_scale is None:
|
||||
@@ -229,12 +228,12 @@ class LocalAttention(nn.Module):
|
||||
dtype,
|
||||
supported_attention_backends=supported_attention_backends)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
causal=causal,
|
||||
**extra_impl_args)
|
||||
self.attn_impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
causal=causal,
|
||||
**extra_impl_args)
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
self.num_kv_heads = num_kv_heads
|
||||
@@ -265,5 +264,5 @@ class LocalAttention(nn.Module):
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
|
||||
output = self.impl.forward(q, k, v, ctx_attn_metadata)
|
||||
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
||||
return output
|
||||
|
||||
@@ -11,13 +11,13 @@ import torch
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.attention.backends.abstract import AttentionBackend
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.platforms import _Backend, current_platform
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum, current_platform
|
||||
from fastvideo.v1.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
|
||||
def backend_name_to_enum(backend_name: str) -> Optional[AttentionBackendEnum]:
|
||||
"""
|
||||
Convert a string backend name to a _Backend enum value.
|
||||
|
||||
@@ -27,11 +27,11 @@ def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
|
||||
loaded.
|
||||
"""
|
||||
assert backend_name is not None
|
||||
return _Backend[backend_name] if backend_name in _Backend.__members__ else \
|
||||
return AttentionBackendEnum[backend_name] if backend_name in AttentionBackendEnum.__members__ else \
|
||||
None
|
||||
|
||||
|
||||
def get_env_variable_attn_backend() -> Optional[_Backend]:
|
||||
def get_env_variable_attn_backend() -> Optional[AttentionBackendEnum]:
|
||||
'''
|
||||
Get the backend override specified by the FastVideo attention
|
||||
backend environment variable, if one is specified.
|
||||
@@ -53,10 +53,11 @@ def get_env_variable_attn_backend() -> Optional[_Backend]:
|
||||
#
|
||||
# THIS SELECTION TAKES PRECEDENCE OVER THE
|
||||
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
|
||||
forced_attn_backend: Optional[_Backend] = None
|
||||
forced_attn_backend: Optional[AttentionBackendEnum] = None
|
||||
|
||||
|
||||
def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
|
||||
def global_force_attn_backend(
|
||||
attn_backend: Optional[AttentionBackendEnum]) -> None:
|
||||
'''
|
||||
Force all attention operations to use a specified backend.
|
||||
|
||||
@@ -71,7 +72,7 @@ def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
|
||||
forced_attn_backend = attn_backend
|
||||
|
||||
|
||||
def get_global_forced_attn_backend() -> Optional[_Backend]:
|
||||
def get_global_forced_attn_backend() -> Optional[AttentionBackendEnum]:
|
||||
'''
|
||||
Get the currently-forced choice of attention backend,
|
||||
or None if auto-selection is currently enabled.
|
||||
@@ -82,7 +83,8 @@ def get_global_forced_attn_backend() -> Optional[_Backend]:
|
||||
def get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
) -> Type[AttentionBackend]:
|
||||
return _cached_get_attn_backend(head_size, dtype,
|
||||
supported_attention_backends)
|
||||
@@ -92,7 +94,8 @@ def get_attn_backend(
|
||||
def _cached_get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
) -> Type[AttentionBackend]:
|
||||
# Check whether a particular choice of backend was
|
||||
# previously forced.
|
||||
@@ -102,7 +105,7 @@ def _cached_get_attn_backend(
|
||||
if not supported_attention_backends:
|
||||
raise ValueError("supported_attention_backends is empty")
|
||||
selected_backend = None
|
||||
backend_by_global_setting: Optional[_Backend] = (
|
||||
backend_by_global_setting: Optional[AttentionBackendEnum] = (
|
||||
get_global_forced_attn_backend())
|
||||
if backend_by_global_setting is not None:
|
||||
selected_backend = backend_by_global_setting
|
||||
@@ -125,7 +128,7 @@ def _cached_get_attn_backend(
|
||||
|
||||
@contextmanager
|
||||
def global_force_attn_backend_context_manager(
|
||||
attn_backend: _Backend) -> Generator[None, None, None]:
|
||||
attn_backend: AttentionBackendEnum) -> Generator[None, None, None]:
|
||||
'''
|
||||
Globally force a FastVideo attention backend override within a
|
||||
context manager, reverting the global attention backend
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field, fields
|
||||
from typing import Any, Dict
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
@@ -12,7 +12,9 @@ logger = init_logger(__name__)
|
||||
# 3. Any field in ArchConfig is fixed upon initialization, and should be hidden away from users
|
||||
@dataclass
|
||||
class ArchConfig:
|
||||
pass
|
||||
stacked_params_mapping: List[Tuple[str, str, str]] = field(
|
||||
default_factory=list
|
||||
) # mapping from huggingface weight names to custom names
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -4,7 +4,7 @@ from typing import Any, List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -12,13 +12,12 @@ class DiTArchConfig(ArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=list)
|
||||
_compile_conditions: list = field(default_factory=list)
|
||||
_param_names_mapping: dict = field(default_factory=dict)
|
||||
_reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
_lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
_supported_attention_backends: Tuple[_Backend,
|
||||
...] = (_Backend.SLIDING_TILE_ATTN,
|
||||
_Backend.SAGE_ATTN,
|
||||
_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA,
|
||||
_Backend.VIDEO_SPARSE_ATTN)
|
||||
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN)
|
||||
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
|
||||
@@ -147,6 +147,9 @@ class HunyuanVideoArchConfig(DiTArchConfig):
|
||||
r"final_layer.linear.\1",
|
||||
})
|
||||
|
||||
# Reverse mapping for saving checkpoints: training -> diffusers
|
||||
_reverse_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
patch_size: int = 2
|
||||
patch_size_t: int = 1
|
||||
in_channels: int = 16
|
||||
|
||||
@@ -5,13 +5,11 @@ from typing import List, Optional, Tuple, Union
|
||||
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
def is_blocks(n: str, m) -> bool:
|
||||
return "blocks" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
@dataclass
|
||||
class StepVideoArchConfig(DiTArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda:
|
||||
[lambda n, m: "transformer_blocks" in n and n.split(".")[-1].isdigit()])
|
||||
|
||||
_param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
|
||||
@@ -49,9 +49,13 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
r"blocks.\1.ffn.fc_in.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
|
||||
r"blocks.\1.ffn.fc_out.\2",
|
||||
r"blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"^blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"blocks.\1.self_attn_residual_norm.norm.\2",
|
||||
})
|
||||
|
||||
# Reverse mapping for saving checkpoints: training -> diffusers
|
||||
_reverse_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
# Some LoRA adapters use the original official layer names instead of hf layer names,
|
||||
# so apply this before the param_names_mapping
|
||||
_lora_param_names_mapping: dict = field(
|
||||
|
||||
@@ -6,14 +6,14 @@ import torch
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
@dataclass
|
||||
class EncoderArchConfig(ArchConfig):
|
||||
architectures: List[str] = field(default_factory=lambda: [])
|
||||
_supported_attention_backends: Tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA)
|
||||
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA)
|
||||
output_hidden_states: bool = False
|
||||
use_return_dict: bool = True
|
||||
|
||||
@@ -32,8 +32,11 @@ class TextEncoderArchConfig(EncoderArchConfig):
|
||||
output_past: bool = True
|
||||
scalable_attention: bool = True
|
||||
tie_word_embeddings: bool = False
|
||||
|
||||
stacked_params_mapping: List[Tuple[str, str, str]] = field(
|
||||
default_factory=list
|
||||
) # mapping from huggingface weight names to custom names
|
||||
tokenizer_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
self.tokenizer_kwargs = {
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
|
||||
ImageEncoderConfig,
|
||||
@@ -8,6 +8,14 @@ from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
|
||||
|
||||
def _is_transformer_layer(n: str, m) -> bool:
|
||||
return "layers" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def _is_embeddings(n: str, m) -> bool:
|
||||
return n.endswith("embeddings")
|
||||
|
||||
|
||||
@dataclass
|
||||
class CLIPTextArchConfig(TextEncoderArchConfig):
|
||||
vocab_size: int = 49408
|
||||
@@ -27,6 +35,15 @@ class CLIPTextArchConfig(TextEncoderArchConfig):
|
||||
bos_token_id: int = 49406
|
||||
eos_token_id: int = 49407
|
||||
text_len: int = 77
|
||||
stacked_params_mapping: List[Tuple[str, str,
|
||||
str]] = field(default_factory=lambda: [
|
||||
# (param_name, shard_name, shard_id)
|
||||
("qkv_proj", "q_proj", "q"),
|
||||
("qkv_proj", "k_proj", "k"),
|
||||
("qkv_proj", "v_proj", "v"),
|
||||
])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda: [_is_transformer_layer, _is_embeddings])
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -1,11 +1,23 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
|
||||
|
||||
def _is_transformer_layer(n: str, m) -> bool:
|
||||
return "layers" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def _is_embeddings(n: str, m) -> bool:
|
||||
return n.endswith("embed_tokens")
|
||||
|
||||
|
||||
def _is_final_norm(n: str, m) -> bool:
|
||||
return n.endswith("norm")
|
||||
|
||||
|
||||
@dataclass
|
||||
class LlamaArchConfig(TextEncoderArchConfig):
|
||||
vocab_size: int = 32000
|
||||
@@ -32,6 +44,18 @@ class LlamaArchConfig(TextEncoderArchConfig):
|
||||
head_dim: Optional[int] = None
|
||||
hidden_state_skip_layer: int = 2
|
||||
text_len: int = 256
|
||||
stacked_params_mapping: List[Tuple[str, str, str]] = field(
|
||||
default_factory=lambda: [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q_proj", "q"),
|
||||
(".qkv_proj", ".k_proj", "k"),
|
||||
(".qkv_proj", ".v_proj", "v"),
|
||||
(".gate_up_proj", ".gate_proj", 0), # type: ignore
|
||||
(".gate_up_proj", ".up_proj", 1), # type: ignore
|
||||
])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda:
|
||||
[_is_transformer_layer, _is_embeddings, _is_final_norm])
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -1,11 +1,23 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
|
||||
TextEncoderConfig)
|
||||
|
||||
|
||||
def _is_transformer_layer(n: str, m) -> bool:
|
||||
return "block" in n and str.isdigit(n.split(".")[-1])
|
||||
|
||||
|
||||
def _is_embeddings(n: str, m) -> bool:
|
||||
return n.endswith("shared")
|
||||
|
||||
|
||||
def _is_final_layernorm(n: str, m) -> bool:
|
||||
return n.endswith("final_layer_norm")
|
||||
|
||||
|
||||
@dataclass
|
||||
class T5ArchConfig(TextEncoderArchConfig):
|
||||
vocab_size: int = 32128
|
||||
@@ -29,6 +41,16 @@ class T5ArchConfig(TextEncoderArchConfig):
|
||||
eos_token_id: int = 1
|
||||
classifier_dropout: float = 0.0
|
||||
text_len: int = 512
|
||||
stacked_params_mapping: List[Tuple[str, str,
|
||||
str]] = field(default_factory=lambda: [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q", "q"),
|
||||
(".qkv_proj", ".k", "k"),
|
||||
(".qkv_proj", ".v", "v"),
|
||||
])
|
||||
_fsdp_shard_conditions: list = field(
|
||||
default_factory=lambda:
|
||||
[_is_transformer_layer, _is_embeddings, _is_final_layernorm])
|
||||
|
||||
# Referenced from https://github.com/huggingface/transformers/blob/main/src/transformers/models/t5/configuration_t5.py
|
||||
def __post_init__(self):
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import argparse
|
||||
import dataclasses
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Union
|
||||
|
||||
@@ -129,3 +131,12 @@ class VAEConfig(ModelConfig):
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "VAEConfig":
|
||||
kwargs = {}
|
||||
for attr in dataclasses.fields(cls):
|
||||
value = getattr(args, attr.name, None)
|
||||
if value is not None:
|
||||
kwargs[attr.name] = value
|
||||
return cls(**kwargs)
|
||||
|
||||
@@ -3,7 +3,7 @@ from fastvideo.v1.configs.pipelines.base import (PipelineConfig,
|
||||
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
|
||||
HunyuanConfig)
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_for_name)
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.v1.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.v1.configs.pipelines.wan import (WanI2V480PConfig,
|
||||
WanI2V720PConfig,
|
||||
@@ -14,5 +14,5 @@ __all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
|
||||
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
|
||||
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
|
||||
"get_pipeline_config_cls_for_name"
|
||||
"get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -1,19 +1,31 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
from dataclasses import asdict, dataclass, field, fields
|
||||
from typing import Any, Callable, Dict, Optional, Tuple, cast
|
||||
from enum import Enum
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
|
||||
VAEConfig)
|
||||
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.v1.configs.utils import update_config_from_args
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import shallow_asdict
|
||||
from fastvideo.v1.utils import (FlexibleArgumentParser, StoreBoolean,
|
||||
shallow_asdict)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class STA_Mode(str, Enum):
|
||||
"""STA (Sliding Tile Attention) modes."""
|
||||
STA_INFERENCE = "STA_inference"
|
||||
STA_SEARCHING = "STA_searching"
|
||||
STA_TUNING = "STA_tuning"
|
||||
STA_TUNING_CFG = "STA_tuning_cfg"
|
||||
NONE = None
|
||||
|
||||
|
||||
def preprocess_text(prompt: str) -> str:
|
||||
return prompt
|
||||
|
||||
@@ -22,59 +34,282 @@ def postprocess_text(output: BaseEncoderOutput) -> torch.tensor:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
# config for a single pipeline
|
||||
@dataclass
|
||||
class PipelineConfig:
|
||||
"""Base configuration for all pipeline architectures."""
|
||||
model_path: str = ""
|
||||
pipeline_config_path: Optional[str] = None
|
||||
|
||||
# Video generation parameters
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: Optional[float] = None
|
||||
disable_autocast: bool = False
|
||||
|
||||
# Model configuration
|
||||
precision: str = "bf16"
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
dit_precision: str = "bf16"
|
||||
|
||||
# VAE configuration
|
||||
vae_config: VAEConfig = field(default_factory=VAEConfig)
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = True
|
||||
vae_config: VAEConfig = field(default_factory=VAEConfig)
|
||||
|
||||
# DiT configuration
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
# Image encoder configuration
|
||||
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp16", ))
|
||||
DEFAULT_TEXT_ENCODER_PRECISIONS = ("fp16", )
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (EncoderConfig(), ))
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp16", ))
|
||||
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (preprocess_text, ))
|
||||
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(postprocess_text, ))
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
# LoRA parameters
|
||||
lora_path: Optional[str] = None
|
||||
lora_nickname: Optional[
|
||||
str] = "default" # for swapping adapters in the pipeline
|
||||
lora_target_names: Optional[List[
|
||||
str]] = None # can restrict list of layers to adapt, e.g. ["q_proj"]
|
||||
|
||||
# StepVideo specific parameters
|
||||
pos_magic: Optional[str] = None
|
||||
neg_magic: Optional[str] = None
|
||||
timesteps_scale: Optional[bool] = None
|
||||
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
STA_mode: str = "STA_inference"
|
||||
STA_mode: STA_Mode = STA_Mode.STA_INFERENCE
|
||||
skip_time_steps: int = 15
|
||||
|
||||
# Compilation
|
||||
enable_torch_compile: bool = False
|
||||
# enable_torch_compile: bool = False
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser,
|
||||
prefix: str = "") -> FlexibleArgumentParser:
|
||||
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
|
||||
|
||||
# model_path will be conflicting with the model_path in FastVideoArgs,
|
||||
# so we add it separately if prefix is not empty
|
||||
if prefix_with_dot != "":
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}model-path",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}model_path",
|
||||
default=PipelineConfig.model_path,
|
||||
help="Path to the pretrained model",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}pipeline-config-path",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}pipeline_config_path",
|
||||
default=PipelineConfig.pipeline_config_path,
|
||||
help="Path to the pipeline config",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}embedded-cfg-scale",
|
||||
type=float,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}embedded_cfg_scale",
|
||||
default=PipelineConfig.embedded_cfg_scale,
|
||||
help="Embedded CFG scale",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}flow-shift",
|
||||
type=float,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}flow_shift",
|
||||
default=PipelineConfig.flow_shift,
|
||||
help="Flow shift parameter",
|
||||
)
|
||||
|
||||
# DiT configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}dit-precision",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}dit_precision",
|
||||
default=PipelineConfig.dit_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for the DiT model",
|
||||
)
|
||||
|
||||
# VAE configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}vae-precision",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}vae_precision",
|
||||
default=PipelineConfig.vae_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for VAE",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}vae-tiling",
|
||||
action=StoreBoolean,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}vae_tiling",
|
||||
default=PipelineConfig.vae_tiling,
|
||||
help="Enable VAE tiling",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}vae-sp",
|
||||
action=StoreBoolean,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}vae_sp",
|
||||
help="Enable VAE spatial parallelism",
|
||||
)
|
||||
|
||||
# Text encoder configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}text-encoder-precisions",
|
||||
nargs="+",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}text_encoder_precisions",
|
||||
default=PipelineConfig.DEFAULT_TEXT_ENCODER_PRECISIONS,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for each text encoder",
|
||||
)
|
||||
|
||||
# Image encoder configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}image-encoder-precision",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}image_encoder_precision",
|
||||
default=PipelineConfig.image_encoder_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for image encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}pos_magic",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}pos_magic",
|
||||
default=PipelineConfig.pos_magic,
|
||||
help="Positive magic prompt for sampling, used in stepvideo",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}neg_magic",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}neg_magic",
|
||||
default=PipelineConfig.neg_magic,
|
||||
help="Negative magic prompt for sampling, used in stepvideo",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}timesteps_scale",
|
||||
type=bool,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}timesteps_scale",
|
||||
default=PipelineConfig.timesteps_scale,
|
||||
help=
|
||||
"Bool for applying scheduler scale in set_timesteps, used in stepvideo",
|
||||
)
|
||||
|
||||
# Add VAE configuration arguments
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEConfig
|
||||
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
|
||||
|
||||
# Add DiT configuration arguments
|
||||
from fastvideo.v1.configs.models.dits.base import DiTConfig
|
||||
DiTConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}dit-config")
|
||||
|
||||
return parser
|
||||
|
||||
def update_config_from_dict(self,
|
||||
args: Dict[str, Any],
|
||||
prefix: str = "") -> None:
|
||||
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
|
||||
update_config_from_args(self, args, prefix, pop_args=True)
|
||||
update_config_from_args(self.vae_config,
|
||||
args,
|
||||
f"{prefix_with_dot}vae_config",
|
||||
pop_args=True)
|
||||
update_config_from_args(self.dit_config,
|
||||
args,
|
||||
f"{prefix_with_dot}dit_config",
|
||||
pop_args=True)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str) -> "PipelineConfig":
|
||||
"""
|
||||
use the pipeline class setting from model_path to match the pipeline config
|
||||
"""
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_for_name)
|
||||
pipeline_config_cls = get_pipeline_config_cls_for_name(model_path)
|
||||
if pipeline_config_cls is not None:
|
||||
pipeline_config = pipeline_config_cls()
|
||||
else:
|
||||
get_pipeline_config_cls_from_name)
|
||||
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
|
||||
|
||||
return cast(PipelineConfig, pipeline_config_cls(model_path=model_path))
|
||||
|
||||
@classmethod
|
||||
def from_kwargs(cls,
|
||||
kwargs: Dict[str, Any],
|
||||
config_cli_prefix: str = "") -> "PipelineConfig":
|
||||
"""
|
||||
Load PipelineConfig from kwargs Dictionary.
|
||||
kwargs: dictionary of kwargs
|
||||
config_cli_prefix: prefix of CLI arguments for this PipelineConfig instance
|
||||
"""
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
|
||||
prefix_with_dot = f"{config_cli_prefix}." if (config_cli_prefix.strip()
|
||||
!= "") else ""
|
||||
model_path: Optional[str] = kwargs.get(prefix_with_dot + 'model_path',
|
||||
None) or kwargs.get('model_path')
|
||||
pipeline_config_or_path: Optional[Union[str, PipelineConfig, Dict[
|
||||
str, Any]]] = kwargs.get(prefix_with_dot + 'pipeline_config',
|
||||
None) or kwargs.get('pipeline_config')
|
||||
if model_path is None:
|
||||
raise ValueError("model_path is required in kwargs")
|
||||
|
||||
# 1. Get the pipeline config class from the registry
|
||||
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
|
||||
|
||||
# 2. Instantiate PipelineConfig
|
||||
if pipeline_config_cls is None:
|
||||
logger.warning(
|
||||
"Couldn't find an optimal sampling param for %s. Using the default sampling param.",
|
||||
"Couldn't find pipeline config for %s. Using the default pipeline config.",
|
||||
model_path)
|
||||
pipeline_config = cls()
|
||||
else:
|
||||
pipeline_config = pipeline_config_cls()
|
||||
|
||||
return cast(PipelineConfig, pipeline_config)
|
||||
# 3. Load PipelineConfig from a json file or a PipelineConfig object if provided
|
||||
if isinstance(pipeline_config_or_path, str):
|
||||
pipeline_config.load_from_json(pipeline_config_or_path)
|
||||
kwargs[prefix_with_dot +
|
||||
'pipeline_config_path'] = pipeline_config_or_path
|
||||
elif isinstance(pipeline_config_or_path, PipelineConfig):
|
||||
pipeline_config = pipeline_config_or_path
|
||||
elif isinstance(pipeline_config_or_path, dict):
|
||||
pipeline_config.update_pipeline_config(pipeline_config_or_path)
|
||||
|
||||
# 4. Update PipelineConfig from CLI arguments if provided
|
||||
kwargs[prefix_with_dot + 'model_path'] = model_path
|
||||
pipeline_config.update_config_from_dict(kwargs, config_cli_prefix)
|
||||
return pipeline_config
|
||||
|
||||
def check_pipeline_config(self) -> None:
|
||||
if self.vae_sp and not self.vae_tiling:
|
||||
raise ValueError(
|
||||
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
|
||||
)
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.text_encoder_precisions):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
|
||||
)
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
)
|
||||
|
||||
if len(self.preprocess_text_funcs) != len(self.postprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
)
|
||||
|
||||
def dump_to_json(self, file_path: str):
|
||||
output_dict = shallow_asdict(self)
|
||||
|
||||
@@ -80,7 +80,7 @@ class HunyuanConfig(PipelineConfig):
|
||||
(llama_postprocess_text, clip_postprocess_text))
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp16", "fp16"))
|
||||
|
||||
@@ -19,7 +19,7 @@ from fastvideo.v1.utils import (maybe_download_model_index,
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Registry maps specific model weights to their config classes
|
||||
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[PipelineConfig]] = {
|
||||
PIPE_NAME_TO_CONFIG: Dict[str, Type[PipelineConfig]] = {
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
|
||||
@@ -51,37 +51,74 @@ PIPELINE_FALLBACK_CONFIG: Dict[str, Type[PipelineConfig]] = {
|
||||
}
|
||||
|
||||
|
||||
def get_pipeline_config_cls_for_name(
|
||||
pipeline_name_or_path: str) -> Optional[type[PipelineConfig]]:
|
||||
"""Get the appropriate config class for specific pretrained weights."""
|
||||
def get_pipeline_config_cls_from_name(
|
||||
pipeline_name_or_path: str) -> Type[PipelineConfig]:
|
||||
"""Get the appropriate configuration class for a given pipeline name or path.
|
||||
|
||||
if os.path.exists(pipeline_name_or_path):
|
||||
config = verify_model_config_and_directory(pipeline_name_or_path)
|
||||
logger.warning(
|
||||
"FastVideo may not correctly identify the optimal config for this model, as the local directory may have been renamed."
|
||||
)
|
||||
else:
|
||||
config = maybe_download_model_index(pipeline_name_or_path)
|
||||
This function implements a multi-step lookup process to find the most suitable
|
||||
configuration class for a given pipeline. It follows this order:
|
||||
1. Exact match in the PIPE_NAME_TO_CONFIG
|
||||
2. Partial match in the PIPE_NAME_TO_CONFIG
|
||||
3. Fallback to class name in the model_index.json
|
||||
4. else raise an error
|
||||
|
||||
pipeline_name = config["_class_name"]
|
||||
Args:
|
||||
pipeline_name_or_path (str): The name or path of the pipeline. This can be:
|
||||
- A registered model ID (e.g., "FastVideo/FastHunyuan-diffusers")
|
||||
- A local path to a model directory
|
||||
- A model ID that will be downloaded
|
||||
|
||||
Returns:
|
||||
Type[PipelineConfig]: The configuration class that best matches the pipeline.
|
||||
This will be one of:
|
||||
- A specific weight configuration class if an exact match is found
|
||||
- A fallback configuration class based on the pipeline architecture
|
||||
- The base PipelineConfig class if no matches are found
|
||||
|
||||
Note:
|
||||
- For local paths, the function will verify the model configuration
|
||||
- For remote models, it will attempt to download the model index
|
||||
- Warning messages are logged when falling back to less specific configurations
|
||||
"""
|
||||
|
||||
pipeline_config_cls: Optional[Type[PipelineConfig]] = None
|
||||
|
||||
# First try exact match for specific weights
|
||||
if pipeline_name_or_path in WEIGHT_CONFIG_REGISTRY:
|
||||
return WEIGHT_CONFIG_REGISTRY[pipeline_name_or_path]
|
||||
if pipeline_name_or_path in PIPE_NAME_TO_CONFIG:
|
||||
pipeline_config_cls = PIPE_NAME_TO_CONFIG[pipeline_name_or_path]
|
||||
|
||||
# Try partial matches (for local paths that might include the weight ID)
|
||||
for registered_id, config_class in WEIGHT_CONFIG_REGISTRY.items():
|
||||
for registered_id, config_class in PIPE_NAME_TO_CONFIG.items():
|
||||
if registered_id in pipeline_name_or_path:
|
||||
return config_class
|
||||
|
||||
# If no match, try to use the fallback config
|
||||
fallback_config = None
|
||||
# Try to determine pipeline architecture for fallback
|
||||
for pipeline_type, detector in PIPELINE_DETECTOR.items():
|
||||
if detector(pipeline_name.lower()):
|
||||
fallback_config = PIPELINE_FALLBACK_CONFIG.get(pipeline_type)
|
||||
pipeline_config_cls = config_class
|
||||
break
|
||||
|
||||
logger.warning("No match found for pipeline %s, using fallback config %s.",
|
||||
pipeline_name_or_path, fallback_config)
|
||||
return fallback_config
|
||||
# If no match, try to use the fallback config
|
||||
if pipeline_config_cls is None:
|
||||
if os.path.exists(pipeline_name_or_path):
|
||||
config = verify_model_config_and_directory(pipeline_name_or_path)
|
||||
else:
|
||||
config = maybe_download_model_index(pipeline_name_or_path)
|
||||
logger.warning(
|
||||
"Trying to use the config from the model_index.json. FastVideo may not correctly identify the optimal config for this model in this situation."
|
||||
)
|
||||
|
||||
pipeline_name = config["_class_name"]
|
||||
# Try to determine pipeline architecture for fallback
|
||||
for pipeline_type, detector in PIPELINE_DETECTOR.items():
|
||||
if detector(pipeline_name.lower()):
|
||||
pipeline_config_cls = PIPELINE_FALLBACK_CONFIG.get(
|
||||
pipeline_type)
|
||||
break
|
||||
|
||||
if pipeline_config_cls is not None:
|
||||
logger.warning(
|
||||
"No match found for pipeline %s, using fallback config %s.",
|
||||
pipeline_name_or_path, pipeline_config_cls)
|
||||
|
||||
if pipeline_config_cls is None:
|
||||
raise ValueError(
|
||||
f"No match found for pipeline {pipeline_name_or_path}, please check the pipeline name or path."
|
||||
)
|
||||
|
||||
return pipeline_config_cls
|
||||
|
||||
@@ -39,7 +39,6 @@ class SamplingParam:
|
||||
num_inference_steps: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
VSA_sparsity: float = 0.0
|
||||
|
||||
# TeaCache parameters
|
||||
enable_teacache: bool = False
|
||||
@@ -185,12 +184,6 @@ class SamplingParam:
|
||||
default=SamplingParam.image_path,
|
||||
help="Path to input image for image-to-video generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--VSA-sparsity",
|
||||
type=float,
|
||||
default=SamplingParam.VSA_sparsity,
|
||||
help="VSA attention sparsity",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
from typing import Any, Dict
|
||||
|
||||
|
||||
def update_config_from_args(config: Any,
|
||||
args_dict: Dict[str, Any],
|
||||
prefix: str = "",
|
||||
pop_args: bool = False) -> None:
|
||||
"""
|
||||
Update configuration object from arguments dictionary.
|
||||
|
||||
Args:
|
||||
config: The configuration object to update
|
||||
args_dict: Dictionary containing arguments
|
||||
prefix: Prefix for the configuration parameters in the args_dict.
|
||||
If None, assumes direct attribute mapping without prefix.
|
||||
"""
|
||||
# Handle top-level attributes (no prefix)
|
||||
args_not_to_remove = [
|
||||
'model_path',
|
||||
]
|
||||
args_to_remove = []
|
||||
if prefix.strip() == "":
|
||||
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)
|
||||
if pop_args:
|
||||
args_to_remove.append(key)
|
||||
else:
|
||||
# Handle nested attributes with prefix
|
||||
prefix_with_dot = f"{prefix}."
|
||||
for key, value in args_dict.items():
|
||||
if key.startswith(prefix_with_dot) and value is not None:
|
||||
attr_name = key[len(prefix_with_dot):]
|
||||
if hasattr(config, attr_name):
|
||||
setattr(config, attr_name, value)
|
||||
if pop_args:
|
||||
args_to_remove.append(key)
|
||||
|
||||
if pop_args:
|
||||
for key in args_to_remove:
|
||||
if key not in args_not_to_remove:
|
||||
args_dict.pop(key)
|
||||
@@ -1,19 +1,17 @@
|
||||
import os
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from torchvision import transforms
|
||||
from torchvision.transforms import Lambda
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.v1.dataset.t2v_datasets import T2V_dataset
|
||||
from fastvideo.v1.dataset.parquet_dataset_map_style import (
|
||||
build_parquet_map_style_dataloader)
|
||||
from fastvideo.v1.dataset.preprocessing_datasets import (
|
||||
VideoCaptionMergedDataset)
|
||||
from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
|
||||
from .parquet_dataset_map_style import build_parquet_map_style_dataloader
|
||||
|
||||
__all__ = ["build_parquet_map_style_dataloader"]
|
||||
from fastvideo.v1.dataset.validation_dataset import ValidationDataset
|
||||
|
||||
|
||||
def getdataset(args, start_idx=0) -> T2V_dataset:
|
||||
def getdataset(args) -> VideoCaptionMergedDataset:
|
||||
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
|
||||
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
|
||||
resize_topcrop = [
|
||||
@@ -31,15 +29,14 @@ def getdataset(args, start_idx=0) -> T2V_dataset:
|
||||
*resize_topcrop,
|
||||
norm_fun,
|
||||
])
|
||||
tokenizer_path = os.path.join(args.model_path, "tokenizer")
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
cache_dir=args.cache_dir)
|
||||
if args.dataset == "t2v":
|
||||
return T2V_dataset(args,
|
||||
transform=transform,
|
||||
temporal_sample=temporal_sample,
|
||||
tokenizer=tokenizer,
|
||||
transform_topcrop=transform_topcrop,
|
||||
start_idx=start_idx)
|
||||
return VideoCaptionMergedDataset(data_merge_path=args.data_merge_path,
|
||||
args=args,
|
||||
transform=transform,
|
||||
temporal_sample=temporal_sample,
|
||||
transform_topcrop=transform_topcrop)
|
||||
|
||||
raise NotImplementedError(args.dataset)
|
||||
|
||||
__all__ = [
|
||||
"build_parquet_map_style_dataloader", "ValidationDataset",
|
||||
"VideoCaptionMergedDataset"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,185 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import argparse
|
||||
import os
|
||||
import pathlib
|
||||
import time
|
||||
|
||||
import torch.distributed as dist
|
||||
import torch.distributed.checkpoint as dist_cp
|
||||
|
||||
from fastvideo.v1.dataset.parquet_dataset_iterable_style import (
|
||||
build_parquet_iterable_style_dataloader)
|
||||
from fastvideo.v1.distributed import get_world_rank
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_local_torch_device,
|
||||
maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Benchmark parquet iterable style dataset loading speed")
|
||||
parser.add_argument(
|
||||
"--path",
|
||||
type=str,
|
||||
help="Path to parquet dataset",
|
||||
)
|
||||
parser.add_argument("--batch_size",
|
||||
type=int,
|
||||
default=4,
|
||||
help="Batch size for DataLoader")
|
||||
parser.add_argument("--num_data_workers",
|
||||
type=int,
|
||||
help="Number of DataLoader workers")
|
||||
parser.add_argument("--num_epoch",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Number of epoches to benchmark")
|
||||
parser.add_argument("--verify_resume",
|
||||
action="store_true",
|
||||
help="Verify resume")
|
||||
parser.add_argument(
|
||||
"--num_batches_per_epoch",
|
||||
type=int,
|
||||
default=1000,
|
||||
help="Number of batches to benchmark",
|
||||
)
|
||||
parser.add_argument('--checkpoint_path',
|
||||
type=str,
|
||||
default='dataloader_checkpoint',
|
||||
help='Path to save/load checkpoint')
|
||||
'''
|
||||
example launch command:
|
||||
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 2 --num_epoch 2 --num_batches_per_epoch 2 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path /mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents/ --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 100
|
||||
'''
|
||||
args = parser.parse_args()
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
maybe_init_distributed_environment_and_model_parallel(
|
||||
tp_size=(world_size + 1) // 2, sp_size=(world_size + 1) // 2)
|
||||
logger.info("Initialized distributed environment with world_size=%d",
|
||||
world_size)
|
||||
|
||||
# Create DataLoader with proper settings
|
||||
dataset, dataloader = build_parquet_iterable_style_dataloader(
|
||||
args.path, args.batch_size, args.num_data_workers)
|
||||
logger.info("Initialized dataloader")
|
||||
|
||||
if args.verify_resume:
|
||||
# First pass - record latent sums
|
||||
first_pass_sums = []
|
||||
for i, (latents, embeddings, masks,
|
||||
caption_text) in enumerate(dataloader):
|
||||
latent_sum = latents.sum().item()
|
||||
first_pass_sums.append(latent_sum)
|
||||
logger.info("Batch %d latent sum: %f", i, latent_sum)
|
||||
if i >= args.num_batches_per_epoch - 1:
|
||||
break
|
||||
|
||||
# Save dataloader state using distributed checkpoint
|
||||
checkpoint_dir = pathlib.Path(args.checkpoint_path)
|
||||
logger.info("Rank %d: Saving dataloader state to %s", get_world_rank(),
|
||||
checkpoint_dir)
|
||||
states = {"dataloader": dataloader}
|
||||
|
||||
begin_time = time.monotonic()
|
||||
dist_cp.save(states, checkpoint_id=checkpoint_dir.as_posix())
|
||||
end_time = time.monotonic()
|
||||
|
||||
logger.info("Rank %d: Saved checkpoint in %.2f seconds",
|
||||
get_world_rank(), end_time - begin_time)
|
||||
|
||||
# Make sure all processes wait for checkpoint to be saved
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
# Recreate dataloader and load state
|
||||
dataset, dataloader = build_parquet_iterable_style_dataloader(
|
||||
args.path, args.batch_size, args.num_data_workers)
|
||||
load_states = {"dataloader": dataloader}
|
||||
dist_cp.load(load_states, checkpoint_id=checkpoint_dir.as_posix())
|
||||
logger.info("Rank %d: Loaded dataloader state from %s",
|
||||
get_world_rank(), checkpoint_dir)
|
||||
|
||||
# Second pass - verify latent sums match
|
||||
for i, (latents, embeddings, masks) in enumerate(dataloader):
|
||||
latent_sum = latents.sum().item()
|
||||
first_pass_sums.append(latent_sum)
|
||||
logger.info("Batch %d latent sum: %f",
|
||||
i + args.num_batches_per_epoch, latent_sum)
|
||||
if i >= args.num_batches_per_epoch - 1:
|
||||
break
|
||||
|
||||
dataset, dataloader = build_parquet_iterable_style_dataloader(
|
||||
args.path, args.batch_size, args.num_data_workers)
|
||||
# Second pass - verify latent sums match
|
||||
second_pass_sums = []
|
||||
for i, (latents, embeddings, masks,
|
||||
caption_text) in enumerate(dataloader):
|
||||
latent_sum = latents.sum().item()
|
||||
second_pass_sums.append(latent_sum)
|
||||
logger.info("Batch %d latent sum: %f (should match first pass: %f)",
|
||||
i, latent_sum, first_pass_sums[i])
|
||||
if i >= args.num_batches_per_epoch * 2 - 1:
|
||||
break
|
||||
|
||||
# Verify all sums match
|
||||
if all(
|
||||
abs(a - b) < 1e-6
|
||||
for a, b in zip(first_pass_sums, second_pass_sums)):
|
||||
logger.info(
|
||||
"All latent sums match between passes - resume verification successful!"
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"Latent sums do not match between passes - resume verification failed!"
|
||||
)
|
||||
|
||||
start_time = time.time()
|
||||
total_samples = 0
|
||||
total_batches = 0
|
||||
for _ in range(args.num_epoch):
|
||||
for i, (latents, embeddings, masks,
|
||||
caption_text) in enumerate(dataloader):
|
||||
if i >= args.num_batches_per_epoch:
|
||||
break
|
||||
|
||||
# Move data to device
|
||||
latents = latents.to(get_local_torch_device())
|
||||
embeddings = embeddings.to(get_local_torch_device())
|
||||
|
||||
# Calculate actual batch size
|
||||
batch_size = latents.size(0)
|
||||
total_samples += batch_size
|
||||
total_batches += 1
|
||||
|
||||
# Print progress only from rank 0
|
||||
if get_world_rank() == 0 and (i + 1) % 10 == 0:
|
||||
elapsed = time.time() - start_time
|
||||
samples_per_sec = total_samples / elapsed
|
||||
logger.info("Batch %d/%d, Speed: %.2f samples/sec", i + 1,
|
||||
args.num_batches_per_epoch, samples_per_sec)
|
||||
|
||||
# Final statistics
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
if get_world_rank() == 0:
|
||||
elapsed = time.time() - start_time
|
||||
samples_per_sec = total_samples / elapsed
|
||||
|
||||
logger.info("\nBenchmark Results:")
|
||||
logger.info("Total time: %.2f seconds", elapsed)
|
||||
logger.info("Total samples: %d", total_samples)
|
||||
logger.info("Average speed: %.2f samples/sec", samples_per_sec)
|
||||
logger.info("Time per batch: %.2f ms", elapsed / total_batches * 1000)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
main()
|
||||
finally:
|
||||
cleanup_dist_env_and_memory()
|
||||
@@ -8,11 +8,12 @@ import torch
|
||||
import torch.distributed as dist
|
||||
import torch.distributed.checkpoint as dist_cp
|
||||
|
||||
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_t2v
|
||||
from fastvideo.v1.dataset.parquet_dataset_map_style import (
|
||||
build_parquet_map_style_dataloader)
|
||||
from fastvideo.v1.distributed import get_world_rank
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_torch_device,
|
||||
cleanup_dist_env_and_memory, get_local_torch_device,
|
||||
maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
@@ -54,9 +55,9 @@ def main() -> None:
|
||||
help='Path to save/load checkpoint')
|
||||
'''
|
||||
example launch command:
|
||||
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 2 --verify_resume
|
||||
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 3 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path /mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents/ --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 100
|
||||
'''
|
||||
args = parser.parse_args()
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
@@ -66,14 +67,22 @@ def main() -> None:
|
||||
world_size)
|
||||
|
||||
# Create DataLoader with proper settings
|
||||
dataloader = build_parquet_map_style_dataloader(args.path, args.batch_size,
|
||||
args.num_data_workers)
|
||||
dataset, dataloader = build_parquet_map_style_dataloader(
|
||||
args.path,
|
||||
args.batch_size,
|
||||
parquet_schema=pyarrow_schema_t2v,
|
||||
num_data_workers=args.num_data_workers)
|
||||
logger.info("Initialized dataloader with %d batches", len(dataloader))
|
||||
|
||||
if args.verify_resume:
|
||||
for i, (latents, embeddings, masks,
|
||||
data_indices) in enumerate(dataloader):
|
||||
logger.info("Batch %d data_indices: %s", i, data_indices)
|
||||
# First pass - record latent sums
|
||||
first_pass_sums = []
|
||||
for i, batch in enumerate(dataloader):
|
||||
latents = batch['vae_latent']
|
||||
embeddings = batch['text_embedding']
|
||||
latent_sum = latents.sum().item()
|
||||
first_pass_sums.append(latent_sum)
|
||||
logger.info("Batch %d latent sum: %f", i, latent_sum)
|
||||
if i >= args.num_batches_per_epoch - 1:
|
||||
break
|
||||
|
||||
@@ -94,47 +103,70 @@ def main() -> None:
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
dataloader = build_parquet_map_style_dataloader(args.path,
|
||||
args.batch_size,
|
||||
args.num_data_workers)
|
||||
# Load dataloader state using distributed checkpoint
|
||||
logger.info("Rank %d: Loading dataloader state from %s",
|
||||
get_world_rank(), checkpoint_dir)
|
||||
# Recreate dataloader and load state
|
||||
dataset, dataloader = build_parquet_map_style_dataloader(
|
||||
args.path,
|
||||
args.batch_size,
|
||||
parquet_schema=pyarrow_schema_t2v,
|
||||
num_data_workers=args.num_data_workers)
|
||||
load_states = {"dataloader": dataloader}
|
||||
dist_cp.load(load_states, checkpoint_id=checkpoint_dir.as_posix())
|
||||
logger.info("Rank %d: Loaded dataloader state from %s",
|
||||
get_world_rank(), checkpoint_dir)
|
||||
|
||||
for i, (latents, embeddings, masks,
|
||||
data_indices) in enumerate(dataloader):
|
||||
logger.info("Batch %d data_indices: %s", i, data_indices)
|
||||
for i, batch in enumerate(dataloader):
|
||||
latents = batch['vae_latent']
|
||||
embeddings = batch['text_embedding']
|
||||
latent_sum = latents.sum().item()
|
||||
first_pass_sums.append(latent_sum)
|
||||
logger.info("Batch %d latent sum: %f",
|
||||
i + args.num_batches_per_epoch, latent_sum)
|
||||
if i >= args.num_batches_per_epoch - 1:
|
||||
break
|
||||
|
||||
logger.info("Restart from the beginning")
|
||||
dataset, dataloader = build_parquet_map_style_dataloader(
|
||||
args.path,
|
||||
args.batch_size,
|
||||
parquet_schema=pyarrow_schema_t2v,
|
||||
num_data_workers=args.num_data_workers)
|
||||
|
||||
dataloader = build_parquet_map_style_dataloader(args.path,
|
||||
args.batch_size,
|
||||
args.num_data_workers)
|
||||
|
||||
for i, (latents, embeddings, masks,
|
||||
data_indices) in enumerate(dataloader):
|
||||
logger.info("Batch %d data_indices: %s", i, data_indices)
|
||||
# Second pass - verify latent sums match
|
||||
second_pass_sums = []
|
||||
for i, batch in enumerate(dataloader):
|
||||
latents = batch['vae_latent']
|
||||
embeddings = batch['text_embedding']
|
||||
latent_sum = latents.sum().item()
|
||||
second_pass_sums.append(latent_sum)
|
||||
logger.info("Batch %d latent sum: %f (should match first pass: %f)",
|
||||
i, latent_sum, first_pass_sums[i])
|
||||
if i >= args.num_batches_per_epoch * 2 - 1:
|
||||
break
|
||||
|
||||
# Verify all sums match
|
||||
if all(
|
||||
abs(a - b) < 1e-6
|
||||
for a, b in zip(first_pass_sums, second_pass_sums)):
|
||||
logger.info(
|
||||
"All latent sums match between passes - resume verification successful!"
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"Latent sums do not match between passes - resume verification failed!"
|
||||
)
|
||||
|
||||
start_time = time.time()
|
||||
total_samples = 0
|
||||
total_batches = 0
|
||||
for _ in range(args.num_epoch):
|
||||
for i, (latents, embeddings, masks,
|
||||
data_indices) in enumerate(dataloader):
|
||||
for i, batch in enumerate(dataloader):
|
||||
latents = batch['vae_latent']
|
||||
embeddings = batch['text_embedding']
|
||||
if i >= args.num_batches_per_epoch:
|
||||
break
|
||||
|
||||
# Move data to device
|
||||
latents = latents.to(get_torch_device())
|
||||
embeddings = embeddings.to(get_torch_device())
|
||||
latents = latents.to(get_local_torch_device())
|
||||
embeddings = embeddings.to(get_local_torch_device())
|
||||
|
||||
# Calculate actual batch size
|
||||
batch_size = latents.size(0)
|
||||
|
||||
@@ -26,15 +26,47 @@ pyarrow_schema_i2v = pa.schema([
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
pa.field("text_attention_mask_bytes", pa.binary()),
|
||||
# e.g., [SeqLen]
|
||||
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bool' or 'int8'
|
||||
pa.field("text_attention_mask_dtype", pa.string()),
|
||||
#I2V
|
||||
pa.field("clip_feature_bytes", pa.binary()),
|
||||
pa.field("clip_feature_shape", pa.list_(pa.int64())),
|
||||
pa.field("clip_feature_dtype", pa.string()),
|
||||
pa.field("first_frame_latent_bytes", pa.binary()),
|
||||
pa.field("first_frame_latent_shape", pa.list_(pa.int64())),
|
||||
pa.field("first_frame_latent_dtype", pa.string()),
|
||||
# I2V Validation
|
||||
pa.field("pil_image_bytes", pa.binary()),
|
||||
pa.field("pil_image_shape", pa.list_(pa.int64())),
|
||||
pa.field("pil_image_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
pa.field("media_type", pa.string()), # 'image' or 'video'
|
||||
pa.field("width", pa.int64()),
|
||||
pa.field("height", pa.int64()),
|
||||
# -- Video-specific (can be null/default for images) ---
|
||||
# Number of frames processed (e.g., 1 for image, N for video)
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
pyarrow_schema_i2v_validation = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
#I2V
|
||||
pa.field("clip_feature_bytes", pa.binary()),
|
||||
pa.field("clip_feature_shape", pa.list_(pa.int64())),
|
||||
pa.field("clip_feature_dtype", pa.string()),
|
||||
# I2V Validation
|
||||
pa.field("pil_image_bytes", pa.binary()),
|
||||
pa.field("pil_image_shape", pa.list_(pa.int64())),
|
||||
pa.field("pil_image_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
@@ -64,11 +96,6 @@ pyarrow_schema_t2v = pa.schema([
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
pa.field("text_attention_mask_bytes", pa.binary()),
|
||||
# e.g., [SeqLen]
|
||||
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bool' or 'int8'
|
||||
pa.field("text_attention_mask_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
@@ -80,4 +107,26 @@ pyarrow_schema_t2v = pa.schema([
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
])
|
||||
|
||||
pyarrow_schema_t2v_validation = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
pa.field("media_type", pa.string()), # 'image' or 'video'
|
||||
pa.field("width", pa.int64()),
|
||||
pa.field("height", pa.int64()),
|
||||
# -- Video-specific (can be null/default for images) ---
|
||||
# Number of frames processed (e.g., 1 for image, N for video)
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
@@ -1,137 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from multiprocessing import Pool, cpu_count
|
||||
from pathlib import Path
|
||||
|
||||
import torchvision
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def get_video_info(video_path):
|
||||
"""Get video information using torchvision."""
|
||||
# Read video tensor (T, C, H, W)
|
||||
video_tensor, _, info = torchvision.io.read_video(str(video_path),
|
||||
output_format="TCHW",
|
||||
pts_unit="sec")
|
||||
|
||||
num_frames = video_tensor.shape[0]
|
||||
height = video_tensor.shape[2]
|
||||
width = video_tensor.shape[3]
|
||||
fps = info.get("video_fps", 0)
|
||||
duration = num_frames / fps if fps > 0 else 0
|
||||
|
||||
# Extract name
|
||||
_, _, videos_dir, video_name = str(video_path).split("/")
|
||||
|
||||
return {
|
||||
"path": str(video_name),
|
||||
"resolution": {
|
||||
"width": width,
|
||||
"height": height
|
||||
},
|
||||
"size": os.path.getsize(video_path),
|
||||
"fps": fps,
|
||||
"duration": duration,
|
||||
"num_frames": num_frames
|
||||
}
|
||||
|
||||
|
||||
def prepare_dataset_json(folder_path,
|
||||
output_name="videos2caption.json",
|
||||
num_workers=None) -> None:
|
||||
"""Prepare dataset information from a folder containing videos and prompt.txt."""
|
||||
folder_path = Path(folder_path)
|
||||
|
||||
# Read prompt file
|
||||
prompt_file = folder_path / "prompt.txt"
|
||||
if not prompt_file.exists():
|
||||
raise FileNotFoundError(f"prompt.txt not found in {folder_path}")
|
||||
|
||||
with open(prompt_file) as f:
|
||||
prompts = [line.strip() for line in f.readlines() if line.strip()]
|
||||
|
||||
# Read videos file
|
||||
videos_file = folder_path / "videos.txt"
|
||||
if not videos_file.exists():
|
||||
raise FileNotFoundError(f"videos.txt not found in {folder_path}")
|
||||
|
||||
with open(videos_file) as f:
|
||||
video_paths = [line.strip() for line in f.readlines() if line.strip()]
|
||||
|
||||
if len(prompts) != len(video_paths):
|
||||
raise ValueError(
|
||||
f"Number of prompts ({len(prompts)}) does not match number of videos ({len(video_paths)})"
|
||||
)
|
||||
|
||||
# Prepare arguments for multiprocessing
|
||||
process_args = [folder_path / video_path for video_path in video_paths]
|
||||
|
||||
# Determine number of workers
|
||||
if num_workers is None:
|
||||
num_workers = max(1, cpu_count() - 1) # Leave one CPU free
|
||||
|
||||
# Process videos in parallel
|
||||
start_time = time.time()
|
||||
with Pool(num_workers) as pool:
|
||||
results = list(
|
||||
tqdm(pool.imap(get_video_info, process_args),
|
||||
total=len(process_args),
|
||||
desc="Processing videos",
|
||||
unit="video"))
|
||||
|
||||
# Combine results with prompts
|
||||
dataset_info = []
|
||||
for result, prompt in zip(results, prompts):
|
||||
result["cap"] = [prompt]
|
||||
dataset_info.append(result)
|
||||
|
||||
# Calculate total processing time
|
||||
total_time = time.time() - start_time
|
||||
total_videos = len(dataset_info)
|
||||
avg_time_per_video = total_time / total_videos if total_videos > 0 else 0
|
||||
|
||||
print("\nProcessing completed:")
|
||||
print(f"Total videos processed: {total_videos}")
|
||||
print(f"Total time: {total_time:.2f} seconds")
|
||||
print(f"Average time per video: {avg_time_per_video:.2f} seconds")
|
||||
|
||||
# Save to JSON file
|
||||
output_file = folder_path / output_name
|
||||
with open(output_file, 'w') as f:
|
||||
json.dump(dataset_info, f, indent=2)
|
||||
|
||||
# Create merge.txt
|
||||
merge_file = folder_path / "merge.txt"
|
||||
with open(merge_file, 'w') as f:
|
||||
f.write(f"{folder_path}/videos,{output_file}\n")
|
||||
|
||||
print(f"Dataset information saved to {output_file}")
|
||||
print(f"Merge file created at {merge_file}")
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description='Prepare video dataset information in JSON format')
|
||||
parser.add_argument(
|
||||
'--folder',
|
||||
type=str,
|
||||
required=True,
|
||||
help='Path to the folder containing videos and prompt.txt')
|
||||
parser.add_argument(
|
||||
'--output',
|
||||
type=str,
|
||||
default='videos2caption.json',
|
||||
help='Name of the output JSON file (default: videos2caption.json)')
|
||||
parser.add_argument('--workers',
|
||||
type=int,
|
||||
default=32,
|
||||
help='Number of worker processes (default: 16)')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = parse_args()
|
||||
prepare_dataset_json(args.folder, args.output, args.workers)
|
||||
@@ -0,0 +1,278 @@
|
||||
import os
|
||||
import pickle
|
||||
import random
|
||||
from typing import Dict, List, Tuple
|
||||
|
||||
import numpy as np
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
import tqdm
|
||||
from torch.utils.data import IterableDataset, get_worker_info
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.v1.dataset.utils import collate_latents_embs_masks
|
||||
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
|
||||
get_world_size)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class BatchIterator:
|
||||
# TODO: Implement state_dict and load_state_dict to support resume.
|
||||
def __init__(self, files, batch_size, text_padding_length, keys,
|
||||
worker_num_samples, read_batch_size):
|
||||
self.files = files
|
||||
self.batch_size = batch_size
|
||||
self.text_padding_length = text_padding_length
|
||||
self.keys = keys
|
||||
self.worker_num_samples = worker_num_samples
|
||||
self.processed_samples = 0
|
||||
self.buffer = []
|
||||
self.read_batch_size = read_batch_size
|
||||
|
||||
def __iter__(self):
|
||||
for file in self.files:
|
||||
if self.processed_samples >= self.worker_num_samples:
|
||||
return
|
||||
|
||||
reader = pq.ParquetFile(file)
|
||||
for batch in reader.iter_batches(batch_size=self.read_batch_size):
|
||||
if self.processed_samples >= self.worker_num_samples:
|
||||
return
|
||||
|
||||
self.buffer.extend(batch.to_pylist())
|
||||
|
||||
while len(self.buffer) >= self.batch_size:
|
||||
if self.processed_samples >= self.worker_num_samples:
|
||||
return
|
||||
|
||||
batch_to_process = self.buffer[:self.batch_size]
|
||||
self.buffer = self.buffer[self.batch_size:]
|
||||
|
||||
all_latents, all_embs, all_masks, caption_text = collate_latents_embs_masks(
|
||||
batch_to_process, self.text_padding_length, self.keys)
|
||||
self.processed_samples += self.batch_size
|
||||
yield all_latents, all_embs, all_masks, caption_text
|
||||
|
||||
|
||||
class LatentsParquetIterStyleDataset(IterableDataset):
|
||||
"""Efficient loader for video-text data from a directory of Parquet files."""
|
||||
|
||||
# Modify this in the future if we want to add more keys, for example, in image to video.
|
||||
keys = [("vae_latent", "latent"), ("text_embedding")]
|
||||
|
||||
def __init__(self,
|
||||
path: str,
|
||||
batch_size: int = 1024,
|
||||
cfg_rate: float = 0.1,
|
||||
num_workers: int = 1,
|
||||
drop_last: bool = True,
|
||||
text_padding_length: int = 512,
|
||||
seed: int = 42,
|
||||
read_batch_size: int = 32,
|
||||
parquet_schema: pa.Schema = None):
|
||||
super().__init__()
|
||||
self.path = str(path)
|
||||
self.batch_size = batch_size
|
||||
self.parquet_schema = parquet_schema
|
||||
self.cfg_rate = cfg_rate
|
||||
self.text_padding_length = text_padding_length
|
||||
self.seed = seed
|
||||
self.read_batch_size = read_batch_size
|
||||
# Get distributed training info
|
||||
self.global_rank = get_world_rank()
|
||||
self.world_size = get_world_size()
|
||||
self.sp_world_size = get_sp_world_size()
|
||||
self.num_sp_groups = self.world_size // self.sp_world_size
|
||||
num_workers = 1 if num_workers == 0 else num_workers
|
||||
# Get sharding info
|
||||
shard_parquet_files, shard_total_samples, shard_parquet_lengths = shard_parquet_files_across_sp_groups_and_workers(
|
||||
self.path, self.num_sp_groups, num_workers, seed)
|
||||
|
||||
if drop_last:
|
||||
self.worker_num_samples = min(
|
||||
shard_total_samples) // batch_size * batch_size
|
||||
# Assign files to current rank's SP group
|
||||
ith_sp_group = self.global_rank // self.sp_world_size
|
||||
self.sp_group_parquet_files = shard_parquet_files[ith_sp_group::self
|
||||
.num_sp_groups]
|
||||
self.sp_group_parquet_lengths = shard_parquet_lengths[
|
||||
ith_sp_group::self.num_sp_groups]
|
||||
self.sp_group_num_samples = shard_total_samples[ith_sp_group::self.
|
||||
num_sp_groups]
|
||||
logger.info(
|
||||
"In total %d parquet files, %d samples, after sharding we retain %d samples due to drop_last",
|
||||
sum([len(shard) for shard in shard_parquet_files]),
|
||||
sum(shard_total_samples),
|
||||
self.worker_num_samples * self.num_sp_groups * num_workers)
|
||||
else:
|
||||
raise ValueError("drop_last must be True")
|
||||
logger.info("Each dataloader worker will load %d samples",
|
||||
self.worker_num_samples)
|
||||
|
||||
def __iter__(self):
|
||||
worker_info = get_worker_info()
|
||||
worker_id = worker_info.id if worker_info is not None else 1
|
||||
|
||||
worker_files = self.sp_group_parquet_files[worker_id]
|
||||
|
||||
batch_iterator = BatchIterator(
|
||||
files=worker_files,
|
||||
batch_size=self.batch_size,
|
||||
text_padding_length=self.text_padding_length,
|
||||
keys=self.keys,
|
||||
worker_num_samples=self.worker_num_samples,
|
||||
read_batch_size=self.read_batch_size) # type: ignore
|
||||
|
||||
yield from batch_iterator
|
||||
|
||||
if batch_iterator.processed_samples != self.worker_num_samples:
|
||||
raise ValueError(
|
||||
"Rank %d, Worker %d: Not enough samples to process, this should not happen",
|
||||
self.global_rank, worker_id)
|
||||
|
||||
|
||||
def shard_parquet_files_across_sp_groups_and_workers(
|
||||
path: str,
|
||||
num_sp_groups: int,
|
||||
num_workers: int,
|
||||
seed: int = 42,
|
||||
) -> Tuple[List[List[str]], List[int], List[Dict[str, int]]]:
|
||||
"""
|
||||
Shard parquet files across SP groups and workers in a balanced way.
|
||||
|
||||
Args:
|
||||
path: Directory containing parquet files
|
||||
num_sp_groups: Number of SP groups to shard across
|
||||
num_workers: Number of workers per SP group
|
||||
seed: Random seed for shuffling
|
||||
|
||||
Returns:
|
||||
Tuple containing:
|
||||
- List of lists of parquet files for each shard
|
||||
- List of total samples per shard
|
||||
- List of dictionaries mapping file paths to their lengths
|
||||
"""
|
||||
# Check if sharding plan already exists
|
||||
sharding_info_dir = os.path.join(
|
||||
path, f"sharding_info_{num_sp_groups}_sp_groups_{num_workers}_workers")
|
||||
if os.path.exists(sharding_info_dir):
|
||||
logger.info("Sharding plan already exists")
|
||||
logger.info("Loading sharding plan from %s", sharding_info_dir)
|
||||
try:
|
||||
with open(
|
||||
os.path.join(sharding_info_dir, "shard_parquet_files.pkl"),
|
||||
"rb") as f:
|
||||
shard_parquet_files = pickle.load(f)
|
||||
with open(
|
||||
os.path.join(sharding_info_dir, "shard_total_samples.pkl"),
|
||||
"rb") as f:
|
||||
shard_total_samples = pickle.load(f)
|
||||
with open(
|
||||
os.path.join(sharding_info_dir,
|
||||
"shard_parquet_lengths.pkl"), "rb") as f:
|
||||
shard_parquet_lengths = pickle.load(f)
|
||||
return shard_parquet_files, shard_total_samples, shard_parquet_lengths
|
||||
except Exception as e:
|
||||
logger.error("Error loading sharding plan: %s", str(e))
|
||||
logger.info("Falling back to creating new sharding plan")
|
||||
|
||||
if get_world_rank() == 0:
|
||||
logger.info("Scanning for parquet files in %s", path)
|
||||
|
||||
# Find all parquet files
|
||||
parquet_files = []
|
||||
|
||||
for root, _, files in os.walk(path):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
parquet_files.append(os.path.join(root, file))
|
||||
|
||||
if not parquet_files:
|
||||
raise ValueError("No parquet files found in %s", path)
|
||||
|
||||
# Calculate file lengths efficiently using a single pass
|
||||
logger.info("Calculating file lengths...")
|
||||
lengths = []
|
||||
for file in tqdm.tqdm(parquet_files, desc="Reading parquet files"):
|
||||
lengths.append(pq.ParquetFile(file).metadata.num_rows)
|
||||
|
||||
total_samples = sum(lengths)
|
||||
logger.info("Found %d files with %d total samples", len(parquet_files),
|
||||
total_samples)
|
||||
|
||||
# Sort files by length for better balancing
|
||||
sorted_indices = np.argsort(lengths)
|
||||
sorted_files = [parquet_files[i] for i in sorted_indices]
|
||||
sorted_lengths = [lengths[i] for i in sorted_indices]
|
||||
|
||||
# Create shards
|
||||
num_shards = num_sp_groups * num_workers
|
||||
shard_parquet_files = [[] for _ in range(num_shards)]
|
||||
shard_total_samples = [0] * num_shards
|
||||
shard_parquet_lengths = [{} for _ in range(num_shards)]
|
||||
|
||||
# Distribute files to shards using a greedy approach
|
||||
logger.info("Distributing files to shards...")
|
||||
for file, length in zip(reversed(sorted_files),
|
||||
reversed(sorted_lengths)):
|
||||
# Find shard with minimum current length
|
||||
target_shard = np.argmin(shard_total_samples)
|
||||
shard_parquet_files[target_shard].append(file)
|
||||
shard_total_samples[target_shard] += length
|
||||
shard_parquet_lengths[target_shard][file] = length
|
||||
#randomize each shard
|
||||
for shard in shard_parquet_files:
|
||||
random.seed(seed)
|
||||
random.shuffle(shard)
|
||||
|
||||
save_dir = os.path.join(
|
||||
path,
|
||||
f"sharding_info_{num_sp_groups}_sp_groups_{num_workers}_workers")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
with open(os.path.join(save_dir, "shard_parquet_files.pkl"), "wb") as f:
|
||||
pickle.dump(shard_parquet_files, f)
|
||||
with open(os.path.join(save_dir, "shard_total_samples.pkl"), "wb") as f:
|
||||
pickle.dump(shard_total_samples, f)
|
||||
with open(os.path.join(save_dir, "shard_parquet_lengths.pkl"),
|
||||
"wb") as f:
|
||||
pickle.dump(shard_parquet_lengths, f)
|
||||
logger.info("Saved sharding info to %s", save_dir)
|
||||
|
||||
# wait for all ranks to finish
|
||||
torch.distributed.barrier()
|
||||
# recursive call
|
||||
return shard_parquet_files_across_sp_groups_and_workers(
|
||||
path, num_sp_groups, num_workers, seed)
|
||||
|
||||
|
||||
def build_parquet_iterable_style_dataloader(
|
||||
path: str,
|
||||
batch_size: int,
|
||||
num_data_workers: int,
|
||||
cfg_rate: float = 0.0,
|
||||
drop_last: bool = True,
|
||||
text_padding_length: int = 512,
|
||||
seed: int = 42,
|
||||
read_batch_size: int = 32
|
||||
) -> Tuple[LatentsParquetIterStyleDataset, StatefulDataLoader]:
|
||||
"""Build a dataloader for the LatentsParquetIterStyleDataset."""
|
||||
dataset = LatentsParquetIterStyleDataset(
|
||||
path=path,
|
||||
batch_size=batch_size,
|
||||
cfg_rate=cfg_rate,
|
||||
num_workers=num_data_workers,
|
||||
drop_last=drop_last,
|
||||
text_padding_length=text_padding_length,
|
||||
seed=seed,
|
||||
read_batch_size=read_batch_size)
|
||||
|
||||
loader = StatefulDataLoader(
|
||||
dataset,
|
||||
batch_size=1,
|
||||
num_workers=num_data_workers,
|
||||
pin_memory=True,
|
||||
)
|
||||
return dataset, loader
|
||||
@@ -1,15 +1,18 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import pickle
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
import numpy as np
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
# Torch in general
|
||||
import torch
|
||||
import tqdm
|
||||
# Dataset
|
||||
from torch.utils.data import Dataset, Sampler
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.v1.dataset.utils import collate_rows_from_parquet_schema
|
||||
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
|
||||
get_world_size)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -30,6 +33,7 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
|
||||
sp_world_size: int,
|
||||
global_rank: int,
|
||||
drop_last: bool = True,
|
||||
drop_first_row: bool = False,
|
||||
seed: int = 0,
|
||||
):
|
||||
self.batch_size = batch_size
|
||||
@@ -45,6 +49,11 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
|
||||
# Create a random permutation of all indices
|
||||
global_indices = torch.randperm(self.dataset_size, generator=rng)
|
||||
|
||||
if drop_first_row:
|
||||
# drop 0 in global_indices
|
||||
global_indices = global_indices[global_indices != 0]
|
||||
self.dataset_size = self.dataset_size - 1
|
||||
|
||||
if self.drop_last:
|
||||
# For drop_last=True, we:
|
||||
# 1. Ensure total samples is divisible by (batch_size * num_sp_groups)
|
||||
@@ -56,19 +65,22 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
|
||||
self.num_sp_groups *
|
||||
self.batch_size]
|
||||
else:
|
||||
# add more indices to make it divisible by (batch_size * num_sp_groups)
|
||||
padding_size = self.num_sp_groups * self.batch_size - (
|
||||
self.dataset_size % (self.num_sp_groups * self.batch_size))
|
||||
global_indices = torch.cat(
|
||||
[global_indices, global_indices[:padding_size]])
|
||||
if self.dataset_size % (self.num_sp_groups * self.batch_size) != 0:
|
||||
# add more indices to make it divisible by (batch_size * num_sp_groups)
|
||||
padding_size = self.num_sp_groups * self.batch_size - (
|
||||
self.dataset_size % (self.num_sp_groups * self.batch_size))
|
||||
logger.info("Padding the dataset from %d to %d",
|
||||
self.dataset_size, self.dataset_size + padding_size)
|
||||
global_indices = torch.cat(
|
||||
[global_indices, global_indices[:padding_size]])
|
||||
|
||||
# shard the indices to each sp group
|
||||
ith_sp_group = self.global_rank // self.sp_world_size
|
||||
sp_group_local_indices = global_indices[ith_sp_group::self.
|
||||
num_sp_groups]
|
||||
|
||||
self.sp_group_local_indices = sp_group_local_indices
|
||||
logger.info("sp_group_local_indices: %d", len(sp_group_local_indices))
|
||||
logger.info("Dataset size for each sp group: %d",
|
||||
len(sp_group_local_indices))
|
||||
|
||||
def __iter__(self):
|
||||
indices = self.sp_group_local_indices
|
||||
@@ -81,20 +93,49 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
|
||||
|
||||
|
||||
def get_parquet_files_and_length(path: str):
|
||||
lengths = []
|
||||
file_names = []
|
||||
for root, _, files in os.walk(path):
|
||||
for file in sorted(files):
|
||||
if file.endswith('.parquet'):
|
||||
file_path = os.path.join(root, file)
|
||||
num_rows = pq.ParquetFile(file_path).metadata.num_rows
|
||||
lengths.append(num_rows)
|
||||
file_names.append(file_path)
|
||||
# sort according to file name to ensure all rank has the same order (in case os.walk is not sorted)
|
||||
file_names_sorted, lengths_sorted = zip(
|
||||
*sorted(zip(file_names, lengths), key=lambda x: x[0]))
|
||||
assert len(file_names_sorted) != 0, "No parquet files found in the dataset"
|
||||
return file_names_sorted, lengths_sorted
|
||||
# Check if cached info exists
|
||||
cache_dir = os.path.join(path, "map_style_cache")
|
||||
cache_file = os.path.join(cache_dir, "file_info.pkl")
|
||||
|
||||
if os.path.exists(cache_file):
|
||||
logger.info("Loading cached file info from %s", cache_file)
|
||||
try:
|
||||
with open(cache_file, "rb") as f:
|
||||
file_names_sorted, lengths_sorted = pickle.load(f)
|
||||
return file_names_sorted, lengths_sorted
|
||||
except Exception as e:
|
||||
logger.error("Error loading cached file info: %s", str(e))
|
||||
logger.info("Falling back to scanning files")
|
||||
|
||||
# If no cache exists or loading failed, scan files
|
||||
if get_world_rank() == 0:
|
||||
lengths = []
|
||||
file_names = []
|
||||
for root, _, files in os.walk(path):
|
||||
for file in sorted(files):
|
||||
if file.endswith('.parquet'):
|
||||
file_path = os.path.join(root, file)
|
||||
file_names.append(file_path)
|
||||
for file_path in tqdm.tqdm(file_names,
|
||||
desc="Reading parquet files to get lengths"):
|
||||
num_rows = pq.ParquetFile(file_path).metadata.num_rows
|
||||
lengths.append(num_rows)
|
||||
# sort according to file name to ensure all rank has the same order (in case os.walk is not sorted)
|
||||
file_names_sorted, lengths_sorted = zip(
|
||||
*sorted(zip(file_names, lengths), key=lambda x: x[0]))
|
||||
assert len(
|
||||
file_names_sorted) != 0, "No parquet files found in the dataset"
|
||||
|
||||
os.makedirs(cache_dir, exist_ok=True)
|
||||
with open(cache_file, "wb") as f:
|
||||
pickle.dump((file_names_sorted, lengths_sorted), f)
|
||||
logger.info("Saved file info to %s", cache_file)
|
||||
|
||||
# Wait for rank 0 to finish saving
|
||||
if get_world_size() > 1:
|
||||
torch.distributed.barrier()
|
||||
|
||||
return get_parquet_files_and_length(path)
|
||||
|
||||
|
||||
def read_row_from_parquet_file(parquet_files: List[str], global_row_idx: int,
|
||||
@@ -144,39 +185,27 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
Note:
|
||||
Using parquet for map style dataset is not efficient, we mainly keep it for backward compatibility and debugging.
|
||||
"""
|
||||
# Modify this in the future if we want to add more keys, for example, in image to video.
|
||||
keys = ["vae_latent", "text_embedding"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path: str,
|
||||
batch_size: int,
|
||||
parquet_schema: pa.Schema,
|
||||
cfg_rate: float = 0.0,
|
||||
seed: int = 42,
|
||||
drop_last: bool = True,
|
||||
drop_first_row: bool = False,
|
||||
text_padding_length: int = 512,
|
||||
):
|
||||
super().__init__()
|
||||
self.path = path
|
||||
self.cfg_rate = cfg_rate
|
||||
if cfg_rate > 0.0:
|
||||
raise ValueError(
|
||||
"cfg_rate > 0.0 is not supported for now because it will trigger bug when num_data_workers > 0"
|
||||
)
|
||||
self.parquet_schema = parquet_schema
|
||||
logger.info("Initializing LatentsParquetMapStyleDataset with path: %s",
|
||||
path)
|
||||
self.parquet_files, self.lengths = get_parquet_files_and_length(path)
|
||||
self.batch = batch_size
|
||||
self.text_padding_length = text_padding_length
|
||||
self._cols = [
|
||||
"vae_latent_bytes",
|
||||
"vae_latent_shape",
|
||||
"text_embedding_bytes",
|
||||
"text_embedding_shape",
|
||||
"text_embedding_dtype",
|
||||
"height",
|
||||
"width",
|
||||
]
|
||||
self.sampler = DP_SP_BatchSampler(
|
||||
batch_size=batch_size,
|
||||
dataset_size=sum(self.lengths),
|
||||
@@ -184,27 +213,14 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
sp_world_size=get_sp_world_size(),
|
||||
global_rank=get_world_rank(),
|
||||
drop_last=drop_last,
|
||||
drop_first_row=drop_first_row,
|
||||
seed=seed,
|
||||
)
|
||||
logger.info("Dataset initialized with %d parquet files and %d rows",
|
||||
len(self.parquet_files), sum(self.lengths))
|
||||
|
||||
def _get_torch_tensors_from_row_dict(
|
||||
self, row_dict: Dict[str, Any]) -> Dict[str, torch.Tensor]:
|
||||
"""
|
||||
Get the latents and prompts from a row dictionary.
|
||||
"""
|
||||
return_dict = {}
|
||||
for key in self.keys:
|
||||
shape = row_dict[f"{key}_shape"]
|
||||
bytes = row_dict[f"{key}_bytes"]
|
||||
# TODO (peiyuan): read precision
|
||||
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
|
||||
data = torch.from_numpy(data)
|
||||
return_dict[key] = data
|
||||
return return_dict
|
||||
|
||||
def get_validation_negative_prompt(self) -> tuple[Any, Any, Any, Any]:
|
||||
def get_validation_negative_prompt(
|
||||
self) -> tuple[torch.Tensor, torch.Tensor, str]:
|
||||
"""
|
||||
Get the negative prompt for validation.
|
||||
This method ensures the negative prompt is loaded and cached properly.
|
||||
@@ -218,39 +234,23 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
row_dict = read_row_from_parquet_file([file_path], row_idx,
|
||||
[self.lengths[0]])
|
||||
|
||||
# Get tensors using the existing helper method
|
||||
data = self._get_torch_tensors_from_row_dict(row_dict)
|
||||
emb = data["text_embedding"]
|
||||
batch = collate_rows_from_parquet_schema([row_dict],
|
||||
self.parquet_schema,
|
||||
self.text_padding_length,
|
||||
cfg_rate=0.0)
|
||||
negative_prompt = batch['info_list'][0]['prompt']
|
||||
negative_prompt_embedding = batch['text_embedding']
|
||||
negative_prompt_attention_mask = batch['text_attention_mask']
|
||||
if len(negative_prompt_embedding.shape) == 2:
|
||||
negative_prompt_embedding = negative_prompt_embedding.unsqueeze(0)
|
||||
if len(negative_prompt_attention_mask.shape) == 1:
|
||||
negative_prompt_attention_mask = negative_prompt_attention_mask.unsqueeze(
|
||||
0).unsqueeze(0)
|
||||
|
||||
# Pad the embedding and get mask
|
||||
padded_emb, mask = self._pad(emb, self.text_padding_length)
|
||||
|
||||
# Pin memory for faster transfer to GPU
|
||||
padded_emb = padded_emb
|
||||
mask = mask
|
||||
|
||||
return None, padded_emb, mask, None
|
||||
|
||||
def _pad(self, t: torch.Tensor, padding_length: int) -> torch.Tensor:
|
||||
"""
|
||||
Pad or crop an embedding [L, D] to exactly padding_length tokens.
|
||||
Return:
|
||||
- [L, D] tensor in pinned CPU memory
|
||||
- [L] attention mask in pinned CPU memory
|
||||
"""
|
||||
L, D = t.shape
|
||||
if padding_length > L: # pad
|
||||
pad = torch.zeros(padding_length - L,
|
||||
D,
|
||||
dtype=t.dtype,
|
||||
device=t.device)
|
||||
return torch.cat([t, pad], 0), torch.cat(
|
||||
[torch.ones(L), torch.zeros(padding_length - L)], 0)
|
||||
else: # crop
|
||||
return t[:padding_length], torch.ones(padding_length)
|
||||
return negative_prompt_embedding, negative_prompt_attention_mask, negative_prompt
|
||||
|
||||
# PyTorch calls this ONLY because the batch_sampler yields a list
|
||||
def __getitems__(self, indices: List[int]):
|
||||
def __getitems__(self, indices: List[int]) -> Dict[str, Any]:
|
||||
"""
|
||||
Batch fetch using read_row_from_parquet_file for each index.
|
||||
"""
|
||||
@@ -259,29 +259,11 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
for idx in indices
|
||||
]
|
||||
|
||||
# Initialize tensors to hold padded embeddings and masks
|
||||
all_latents = []
|
||||
all_embs = []
|
||||
all_masks = []
|
||||
|
||||
# Process each row individually
|
||||
for i, row in enumerate(rows):
|
||||
# Get tensors from row
|
||||
data = self._get_torch_tensors_from_row_dict(row)
|
||||
latents, emb = data["vae_latent"], data["text_embedding"]
|
||||
|
||||
padded_emb, mask = self._pad(emb, self.text_padding_length)
|
||||
# Store in batch tensors
|
||||
all_latents.append(latents)
|
||||
all_embs.append(padded_emb)
|
||||
all_masks.append(mask)
|
||||
|
||||
# Pin memory for faster transfer to GPU
|
||||
all_latents = torch.stack(all_latents)
|
||||
all_embs = torch.stack(all_embs)
|
||||
all_masks = torch.stack(all_masks)
|
||||
|
||||
return all_latents, all_embs, all_masks, indices
|
||||
batch = collate_rows_from_parquet_schema(rows,
|
||||
self.parquet_schema,
|
||||
self.text_padding_length,
|
||||
cfg_rate=self.cfg_rate)
|
||||
return batch
|
||||
|
||||
def __len__(self):
|
||||
return sum(self.lengths)
|
||||
@@ -298,8 +280,10 @@ def build_parquet_map_style_dataloader(
|
||||
path,
|
||||
batch_size,
|
||||
num_data_workers,
|
||||
parquet_schema,
|
||||
cfg_rate=0.0,
|
||||
drop_last=True,
|
||||
drop_first_row=False,
|
||||
text_padding_length=512,
|
||||
seed=42) -> Tuple[LatentsParquetMapStyleDataset, StatefulDataLoader]:
|
||||
dataset = LatentsParquetMapStyleDataset(
|
||||
@@ -307,7 +291,9 @@ def build_parquet_map_style_dataloader(
|
||||
batch_size,
|
||||
cfg_rate=cfg_rate,
|
||||
drop_last=drop_last,
|
||||
drop_first_row=drop_first_row,
|
||||
text_padding_length=text_padding_length,
|
||||
parquet_schema=parquet_schema,
|
||||
seed=seed)
|
||||
|
||||
loader = StatefulDataLoader(
|
||||
|
||||
@@ -0,0 +1,615 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import Counter
|
||||
from dataclasses import dataclass
|
||||
from os.path import join as opj
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class PreprocessBatch:
|
||||
"""
|
||||
Batch information for dataset processing stages.
|
||||
|
||||
This class holds all the information about a video-caption or image-caption pair
|
||||
as it moves through the processing pipeline. Fields are populated by different stages.
|
||||
"""
|
||||
# Raw metadata
|
||||
path: str
|
||||
cap: Union[str, List[str]]
|
||||
resolution: Optional[Dict] = None
|
||||
fps: Optional[float] = None
|
||||
duration: Optional[float] = None
|
||||
|
||||
# Processed metadata
|
||||
num_frames: Optional[int] = None
|
||||
sample_frame_index: Optional[List[int]] = None
|
||||
sample_num_frames: Optional[int] = None
|
||||
|
||||
# Processed data
|
||||
pixel_values: Optional[torch.Tensor] = None
|
||||
text: Optional[str] = None
|
||||
input_ids: Optional[torch.Tensor] = None
|
||||
cond_mask: Optional[torch.Tensor] = None
|
||||
|
||||
@property
|
||||
def is_video(self) -> bool:
|
||||
"""Check if this is a video item."""
|
||||
return self.path.endswith(".mp4")
|
||||
|
||||
@property
|
||||
def is_image(self) -> bool:
|
||||
"""Check if this is an image item."""
|
||||
return self.path.endswith(".jpg")
|
||||
|
||||
|
||||
class DatasetStage(ABC):
|
||||
"""
|
||||
Abstract base class for dataset processing stages.
|
||||
|
||||
Similar to PipelineStage but designed for dataset preprocessing operations.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Process the dataset batch.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch to process
|
||||
**kwargs: Additional processing parameters
|
||||
|
||||
Returns:
|
||||
Processed batch
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class DatasetFilterStage(ABC):
|
||||
"""
|
||||
Abstract base class for dataset filtering stages.
|
||||
|
||||
These stages can filter out items during metadata processing.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
|
||||
"""
|
||||
Check if batch should be kept.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch to check
|
||||
**kwargs: Additional parameters
|
||||
|
||||
Returns:
|
||||
True if batch should be kept, False otherwise
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Process the dataset batch (for non-filtering operations).
|
||||
|
||||
Args:
|
||||
batch: Dataset batch to process
|
||||
**kwargs: Additional processing parameters
|
||||
|
||||
Returns:
|
||||
Processed batch
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class DataValidationStage(DatasetFilterStage):
|
||||
"""Stage for validating data items."""
|
||||
|
||||
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
|
||||
"""
|
||||
Validate data item.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch to validate
|
||||
|
||||
Returns:
|
||||
True if valid, False if invalid
|
||||
"""
|
||||
# Check for caption
|
||||
if batch.cap is None:
|
||||
return False
|
||||
|
||||
if batch.is_video:
|
||||
# Validate video-specific fields
|
||||
if batch.duration is None or batch.fps is None:
|
||||
return False
|
||||
elif not batch.is_image:
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""Process does nothing for validation - filtering is handled by should_keep."""
|
||||
return batch
|
||||
|
||||
|
||||
class ResolutionFilterStage(DatasetFilterStage):
|
||||
"""Stage for filtering data items based on resolution constraints."""
|
||||
|
||||
def __init__(self,
|
||||
max_h_div_w_ratio: float = 17 / 16,
|
||||
min_h_div_w_ratio: float = 8 / 16,
|
||||
max_height: int = 1024,
|
||||
max_width: int = 1024):
|
||||
self.max_h_div_w_ratio = max_h_div_w_ratio
|
||||
self.min_h_div_w_ratio = min_h_div_w_ratio
|
||||
self.max_height = max_height
|
||||
self.max_width = max_width
|
||||
|
||||
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
|
||||
"""
|
||||
Check if data item passes resolution filtering.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch with resolution information
|
||||
|
||||
Returns:
|
||||
True if passes filter, False otherwise
|
||||
"""
|
||||
# Only apply to videos
|
||||
if not batch.is_video:
|
||||
return True
|
||||
|
||||
if batch.resolution is None:
|
||||
return False
|
||||
|
||||
height = batch.resolution.get("height", None)
|
||||
width = batch.resolution.get("width", None)
|
||||
if height is None or width is None:
|
||||
return False
|
||||
|
||||
# Check aspect ratio
|
||||
aspect = self.max_height / self.max_width
|
||||
hw_aspect_thr = 1.5
|
||||
|
||||
return self.filter_resolution(
|
||||
height,
|
||||
width,
|
||||
max_h_div_w_ratio=hw_aspect_thr * aspect,
|
||||
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
|
||||
)
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""Process does nothing for resolution filtering - filtering is handled by should_keep."""
|
||||
return batch
|
||||
|
||||
def filter_resolution(self, h: int, w: int, max_h_div_w_ratio: float,
|
||||
min_h_div_w_ratio: float) -> bool:
|
||||
"""Filter based on height/width ratio."""
|
||||
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
|
||||
|
||||
|
||||
class FrameSamplingStage(DatasetFilterStage):
|
||||
"""Stage for temporal frame sampling and indexing."""
|
||||
|
||||
def __init__(self,
|
||||
num_frames: int,
|
||||
train_fps: int,
|
||||
speed_factor: int = 1,
|
||||
video_length_tolerance_range: float = 5.0,
|
||||
drop_short_ratio: float = 0.0):
|
||||
self.num_frames = num_frames
|
||||
self.train_fps = train_fps
|
||||
self.speed_factor = speed_factor
|
||||
self.video_length_tolerance_range = video_length_tolerance_range
|
||||
self.drop_short_ratio = drop_short_ratio
|
||||
|
||||
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
|
||||
"""
|
||||
Check if video should be kept based on length constraints.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch
|
||||
|
||||
Returns:
|
||||
True if should be kept, False otherwise
|
||||
"""
|
||||
if batch.is_image:
|
||||
return True
|
||||
|
||||
if batch.duration is None or batch.fps is None:
|
||||
return False
|
||||
|
||||
num_frames = math.ceil(batch.fps * batch.duration)
|
||||
|
||||
# Check if video is too long
|
||||
if (num_frames / batch.fps > self.video_length_tolerance_range *
|
||||
(self.num_frames / self.train_fps * self.speed_factor)):
|
||||
return False
|
||||
|
||||
# Resample frame indices to check length
|
||||
frame_interval = batch.fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, num_frames,
|
||||
frame_interval).astype(int)
|
||||
|
||||
# Filter short videos
|
||||
return not (len(frame_indices) < self.num_frames
|
||||
and random.random() < self.drop_short_ratio)
|
||||
|
||||
def process(self,
|
||||
batch: PreprocessBatch,
|
||||
temporal_sample_fn=None,
|
||||
**kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Process frame sampling for video data items.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch
|
||||
temporal_sample_fn: Function for temporal sampling
|
||||
|
||||
Returns:
|
||||
Updated batch with frame sampling info
|
||||
"""
|
||||
if batch.is_image:
|
||||
# For images, just add sample info
|
||||
batch.sample_frame_index = [0]
|
||||
batch.sample_num_frames = 1
|
||||
return batch
|
||||
|
||||
assert batch.duration is not None and batch.fps is not None
|
||||
batch.num_frames = math.ceil(batch.fps * batch.duration)
|
||||
|
||||
# Resample frame indices
|
||||
frame_interval = batch.fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, batch.num_frames,
|
||||
frame_interval).astype(int)
|
||||
|
||||
# Temporal crop if too long
|
||||
if len(frame_indices
|
||||
) > self.num_frames and temporal_sample_fn is not None:
|
||||
begin_index, end_index = temporal_sample_fn(len(frame_indices))
|
||||
frame_indices = frame_indices[begin_index:end_index]
|
||||
|
||||
batch.sample_frame_index = frame_indices.tolist()
|
||||
batch.sample_num_frames = len(frame_indices)
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
class VideoTransformStage(DatasetStage):
|
||||
"""Stage for video data transformation."""
|
||||
|
||||
def __init__(self, transform) -> None:
|
||||
self.transform = transform
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Transform video data.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch with video information
|
||||
|
||||
Returns:
|
||||
Batch with transformed video tensor
|
||||
"""
|
||||
if not batch.is_video:
|
||||
return batch
|
||||
|
||||
assert os.path.exists(batch.path), f"file {batch.path} do not exist!"
|
||||
assert batch.sample_frame_index is not None, "Frame indices must be set before transformation"
|
||||
|
||||
torchvision_video, _, metadata = torchvision.io.read_video(
|
||||
batch.path, output_format="TCHW")
|
||||
video = torchvision_video[batch.sample_frame_index]
|
||||
if self.transform is not None:
|
||||
video = self.transform(video)
|
||||
video = rearrange(video, "t c h w -> c t h w")
|
||||
video = video.to(torch.uint8)
|
||||
|
||||
h, w = video.shape[-2:]
|
||||
assert (
|
||||
h / w <= 17 / 16 and h / w >= 8 / 16
|
||||
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({batch.path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
|
||||
|
||||
video = video.float() / 127.5 - 1.0
|
||||
batch.pixel_values = video
|
||||
return batch
|
||||
|
||||
|
||||
class ImageTransformStage(DatasetStage):
|
||||
"""Stage for image data transformation."""
|
||||
|
||||
def __init__(self, transform, transform_topcrop) -> None:
|
||||
self.transform = transform
|
||||
self.transform_topcrop = transform_topcrop
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Transform image data.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch with image information
|
||||
|
||||
Returns:
|
||||
Batch with transformed image tensor
|
||||
"""
|
||||
if not batch.is_image:
|
||||
return batch
|
||||
|
||||
image = Image.open(batch.path).convert("RGB")
|
||||
image = torch.from_numpy(np.array(image))
|
||||
image = rearrange(image, "h w c -> c h w").unsqueeze(0)
|
||||
|
||||
if self.transform_topcrop is not None:
|
||||
image = self.transform_topcrop(image)
|
||||
elif self.transform is not None:
|
||||
image = self.transform(image)
|
||||
|
||||
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
|
||||
image = image.float() / 127.5 - 1.0
|
||||
batch.pixel_values = image
|
||||
return batch
|
||||
|
||||
|
||||
class TextEncodingStage(DatasetStage):
|
||||
"""Stage for text tokenization and encoding."""
|
||||
|
||||
def __init__(self, tokenizer, text_max_length: int, cfg_rate: float = 0.0):
|
||||
self.tokenizer = tokenizer
|
||||
self.text_max_length = text_max_length
|
||||
self.cfg_rate = cfg_rate
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Process text data.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch with caption information
|
||||
|
||||
Returns:
|
||||
Batch with encoded text information
|
||||
"""
|
||||
text = batch.cap
|
||||
if not isinstance(text, list):
|
||||
text = [text]
|
||||
text = [random.choice(text)]
|
||||
|
||||
text = text[0] if random.random() > self.cfg_rate else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
batch.text = text
|
||||
batch.input_ids = text_tokens_and_mask["input_ids"]
|
||||
batch.cond_mask = text_tokens_and_mask["attention_mask"]
|
||||
return batch
|
||||
|
||||
|
||||
class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
torch.distributed.checkpoint.stateful.Stateful):
|
||||
"""
|
||||
Merged dataset for video and caption data with stage-based processing.
|
||||
Assumes that data_merge_path is a txt file with the following format:
|
||||
<folder_path>,<json_file_path>
|
||||
|
||||
The folder should contain videos.
|
||||
|
||||
The json file should be a list of dictionaries with the following format:
|
||||
[
|
||||
{
|
||||
"path": "1gGQy4nxyUo-Scene-016.mp4",
|
||||
"resolution": {
|
||||
"width": 1920,
|
||||
"height": 1080
|
||||
},
|
||||
"size": 2439112,
|
||||
"fps": 25.0,
|
||||
"duration": 6.88,
|
||||
"num_frames": 172,
|
||||
"cap": [
|
||||
"A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open."
|
||||
]
|
||||
},
|
||||
...
|
||||
]
|
||||
|
||||
This dataset processes video and image data through a series of stages:
|
||||
- Data validation
|
||||
- Resolution filtering
|
||||
- Frame sampling
|
||||
- Transformation
|
||||
- Text encoding
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
data_merge_path: str,
|
||||
args,
|
||||
transform,
|
||||
temporal_sample,
|
||||
transform_topcrop,
|
||||
start_idx: int = 0):
|
||||
self.data_merge_path = data_merge_path
|
||||
self.start_idx = start_idx
|
||||
self.args = args
|
||||
self.temporal_sample = temporal_sample
|
||||
|
||||
# Initialize tokenizer
|
||||
tokenizer_path = os.path.join(args.model_path, "tokenizer")
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
cache_dir=args.cache_dir)
|
||||
|
||||
# Initialize processing stages
|
||||
self._init_stages(args, transform, transform_topcrop, tokenizer)
|
||||
|
||||
# Process metadata
|
||||
self.processed_batches = self._process_metadata()
|
||||
|
||||
def _init_stages(self, args, transform, transform_topcrop,
|
||||
tokenizer) -> None:
|
||||
"""Initialize all processing stages."""
|
||||
self.validation_stage = DataValidationStage()
|
||||
self.resolution_filter_stage = ResolutionFilterStage(
|
||||
max_height=args.max_height, max_width=args.max_width)
|
||||
self.frame_sampling_stage = FrameSamplingStage(
|
||||
num_frames=args.num_frames,
|
||||
train_fps=args.train_fps,
|
||||
speed_factor=args.speed_factor,
|
||||
video_length_tolerance_range=args.video_length_tolerance_range,
|
||||
drop_short_ratio=args.drop_short_ratio)
|
||||
self.video_transform_stage = VideoTransformStage(transform)
|
||||
self.image_transform_stage = ImageTransformStage(
|
||||
transform, transform_topcrop)
|
||||
self.text_encoding_stage = TextEncodingStage(
|
||||
tokenizer=tokenizer,
|
||||
text_max_length=args.text_max_length,
|
||||
cfg_rate=args.training_cfg_rate)
|
||||
|
||||
def _load_raw_data(self) -> List[Dict]:
|
||||
"""Load raw data from JSON files."""
|
||||
# Read folder-annotation pairs
|
||||
with open(self.data_merge_path) as f:
|
||||
folder_anno_pairs = [
|
||||
line.strip().split(",") for line in f if line.strip()
|
||||
]
|
||||
assert len(
|
||||
folder_anno_pairs) == 1, "Only support one folder-annotation pair"
|
||||
assert len(folder_anno_pairs[0]
|
||||
) == 2, "Folder-annotation pair should have two elements"
|
||||
folder, annotation_file = folder_anno_pairs[0]
|
||||
|
||||
data_items: List[Dict] = []
|
||||
with open(annotation_file) as f:
|
||||
data_items = json.load(f)
|
||||
|
||||
# Update paths with folder prefix
|
||||
for item in data_items:
|
||||
item["path"] = opj(folder, item["path"])
|
||||
|
||||
return data_items
|
||||
|
||||
def _process_metadata(self) -> List[PreprocessBatch]:
|
||||
"""Process the raw metadata through all filtering stages."""
|
||||
raw_data = self._load_raw_data()
|
||||
processed_batches = []
|
||||
|
||||
# Initialize counters
|
||||
filter_counts = {
|
||||
"validation_failed": 0,
|
||||
"resolution_failed": 0,
|
||||
"frame_sampling_failed": 0
|
||||
}
|
||||
sample_num_frames: List[int] = []
|
||||
|
||||
for item in raw_data:
|
||||
batch = PreprocessBatch(path=item["path"],
|
||||
cap=item["cap"],
|
||||
resolution=item.get("resolution"),
|
||||
fps=item.get("fps"),
|
||||
duration=item.get("duration"))
|
||||
|
||||
# Apply filtering stages
|
||||
if not self._apply_filter_stages(batch, filter_counts):
|
||||
continue
|
||||
|
||||
# Apply frame sampling processing
|
||||
batch = self.frame_sampling_stage.process(
|
||||
batch, temporal_sample_fn=self.temporal_sample)
|
||||
|
||||
processed_batches.append(batch)
|
||||
assert batch.sample_num_frames is not None
|
||||
sample_num_frames.append(batch.sample_num_frames)
|
||||
|
||||
self._log_filtering_stats(filter_counts, sample_num_frames,
|
||||
len(raw_data), len(processed_batches))
|
||||
return processed_batches
|
||||
|
||||
def _apply_filter_stages(self, batch: PreprocessBatch,
|
||||
filter_counts: Dict[str, int]) -> bool:
|
||||
"""Apply all filter stages and update counters. Returns True if batch should be kept."""
|
||||
if not self.validation_stage.should_keep(batch):
|
||||
filter_counts["validation_failed"] += 1
|
||||
return False
|
||||
|
||||
if not self.resolution_filter_stage.should_keep(batch):
|
||||
filter_counts["resolution_failed"] += 1
|
||||
return False
|
||||
|
||||
if not self.frame_sampling_stage.should_keep(batch):
|
||||
filter_counts["frame_sampling_failed"] += 1
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def _log_filtering_stats(self, filter_counts: Dict[str, int],
|
||||
sample_num_frames: List[int], before_count: int,
|
||||
after_count: int):
|
||||
"""Log filtering statistics."""
|
||||
logger.info(
|
||||
"validation_failed: %d, resolution_failed: %d, frame_sampling_failed: %d, "
|
||||
"Counter(sample_num_frames): %s, before filter: %d, after filter: %d",
|
||||
filter_counts['validation_failed'],
|
||||
filter_counts['resolution_failed'],
|
||||
filter_counts['frame_sampling_failed'], Counter(sample_num_frames),
|
||||
before_count, after_count)
|
||||
|
||||
def __iter__(self):
|
||||
"""Iterate through processed data items."""
|
||||
for idx in range(len(self.processed_batches)):
|
||||
yield self._get_item(idx)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.processed_batches)
|
||||
|
||||
def _get_item(self, idx: int) -> Dict:
|
||||
"""Get a single processed data item."""
|
||||
batch = self.processed_batches[idx]
|
||||
|
||||
# Apply transformation stages
|
||||
batch = self.video_transform_stage.process(batch)
|
||||
batch = self.image_transform_stage.process(batch)
|
||||
batch = self.text_encoding_stage.process(batch)
|
||||
|
||||
# Build result dictionary
|
||||
result = {
|
||||
"pixel_values": batch.pixel_values,
|
||||
"text": batch.text,
|
||||
"input_ids": batch.input_ids,
|
||||
"cond_mask": batch.cond_mask,
|
||||
"path": batch.path,
|
||||
}
|
||||
|
||||
# Add video-specific fields
|
||||
if batch.is_video:
|
||||
result.update({"fps": batch.fps, "duration": batch.duration})
|
||||
|
||||
return result
|
||||
|
||||
def state_dict(self) -> Dict[str, Any]:
|
||||
"""Return state dict for checkpointing."""
|
||||
return {"processed_batches": self.processed_batches}
|
||||
|
||||
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
|
||||
"""Load state dict from checkpoint."""
|
||||
self.processed_batches = state_dict["processed_batches"]
|
||||
@@ -1,352 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from collections import Counter
|
||||
from os.path import join as opj
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from fastvideo.utils.dataset_utils import DecordInit
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
|
||||
|
||||
class SingletonMeta(type):
|
||||
_instances: dict[type, 'SingletonMeta'] = {}
|
||||
|
||||
def __call__(cls, *args, **kwargs):
|
||||
if cls not in cls._instances:
|
||||
instance = super().__call__(*args, **kwargs)
|
||||
cls._instances[cls] = instance
|
||||
return cls._instances[cls]
|
||||
|
||||
|
||||
class DataSetProg(metaclass=SingletonMeta):
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.cap_list: list[dict] = []
|
||||
self.elements: list[int] = []
|
||||
self.num_workers = 1
|
||||
self.n_elements = 0
|
||||
self.worker_elements: dict[int, list[int]] = {}
|
||||
self.n_used_elements: dict[int, int] = {}
|
||||
|
||||
def set_cap_list(self, num_workers, cap_list, n_elements) -> None:
|
||||
self.num_workers = num_workers
|
||||
self.cap_list = cap_list
|
||||
self.n_elements = n_elements
|
||||
self.elements = list(range(n_elements))
|
||||
random.shuffle(self.elements)
|
||||
print(f"n_elements: {len(self.elements)}", flush=True)
|
||||
|
||||
for i in range(self.num_workers):
|
||||
self.n_used_elements[i] = 0
|
||||
per_worker = int(
|
||||
math.ceil(len(self.elements) / float(self.num_workers)))
|
||||
start = i * per_worker
|
||||
end = min(start + per_worker, len(self.elements))
|
||||
self.worker_elements[i] = self.elements[start:end]
|
||||
|
||||
def get_item(self, work_info) -> int:
|
||||
worker_id = 0 if work_info is None else work_info.id
|
||||
|
||||
idx = self.worker_elements[worker_id][
|
||||
self.n_used_elements[worker_id] %
|
||||
len(self.worker_elements[worker_id])]
|
||||
self.n_used_elements[worker_id] += 1
|
||||
return idx
|
||||
|
||||
|
||||
dataset_prog = DataSetProg()
|
||||
|
||||
|
||||
def filter_resolution(h: int,
|
||||
w: int,
|
||||
max_h_div_w_ratio: float = 17 / 16,
|
||||
min_h_div_w_ratio: float = 8 / 16) -> bool:
|
||||
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
|
||||
|
||||
|
||||
class T2V_dataset(Dataset):
|
||||
|
||||
def __init__(self,
|
||||
args,
|
||||
transform,
|
||||
temporal_sample,
|
||||
tokenizer,
|
||||
transform_topcrop,
|
||||
start_idx=0) -> None:
|
||||
self.start_idx = start_idx
|
||||
self.data = args.data_merge_path
|
||||
self.num_frames = args.num_frames
|
||||
self.train_fps = args.train_fps
|
||||
self.use_image_num = args.use_image_num
|
||||
self.transform = transform
|
||||
self.transform_topcrop = transform_topcrop
|
||||
self.temporal_sample = temporal_sample
|
||||
self.tokenizer = tokenizer
|
||||
self.text_max_length = args.text_max_length
|
||||
self.cfg = args.cfg
|
||||
self.speed_factor = args.speed_factor
|
||||
self.max_height = args.max_height
|
||||
self.max_width = args.max_width
|
||||
self.drop_short_ratio = args.drop_short_ratio
|
||||
assert self.speed_factor >= 1
|
||||
self.v_decoder = DecordInit()
|
||||
self.video_length_tolerance_range = args.video_length_tolerance_range
|
||||
self.support_Chinese = True
|
||||
if "mt5" not in args.text_encoder_name:
|
||||
self.support_Chinese = False
|
||||
|
||||
cap_list = self.get_cap_list()
|
||||
|
||||
assert len(cap_list) > 0
|
||||
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
|
||||
self.lengths = self.sample_num_frames
|
||||
|
||||
n_elements = len(cap_list)
|
||||
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list,
|
||||
n_elements)
|
||||
|
||||
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
|
||||
|
||||
def set_checkpoint(self, n_used_elements):
|
||||
for i in range(len(dataset_prog.n_used_elements)):
|
||||
dataset_prog.n_used_elements[i] = n_used_elements
|
||||
|
||||
def __len__(self):
|
||||
return dataset_prog.n_elements
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
data = self.get_data(idx)
|
||||
return data
|
||||
|
||||
def get_data(self, idx) -> dict:
|
||||
path = dataset_prog.cap_list[idx]["path"]
|
||||
if path.endswith(".mp4"):
|
||||
return self.get_video(idx)
|
||||
else:
|
||||
return self.get_image(idx)
|
||||
|
||||
def get_video(self, idx) -> dict:
|
||||
video_path = dataset_prog.cap_list[idx]["path"]
|
||||
assert os.path.exists(video_path), f"file {video_path} do not exist!"
|
||||
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
|
||||
|
||||
torchvision_video, _, metadata = torchvision.io.read_video(
|
||||
video_path, output_format="TCHW")
|
||||
video = torchvision_video[frame_indices]
|
||||
video = self.transform(video)
|
||||
video = rearrange(video, "t c h w -> c t h w")
|
||||
video = video.to(torch.uint8)
|
||||
assert video.dtype == torch.uint8
|
||||
|
||||
h, w = video.shape[-2:]
|
||||
assert (
|
||||
h / w <= 17 / 16 and h / w >= 8 / 16
|
||||
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
|
||||
|
||||
video = video.float() / 127.5 - 1.0
|
||||
|
||||
text = dataset_prog.cap_list[idx]["cap"]
|
||||
if not isinstance(text, list):
|
||||
text = [text]
|
||||
text = [random.choice(text)]
|
||||
|
||||
text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"]
|
||||
cond_mask = text_tokens_and_mask["attention_mask"]
|
||||
return dict(pixel_values=video,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=video_path,
|
||||
fps=dataset_prog.cap_list[idx]["fps"],
|
||||
duration=dataset_prog.cap_list[idx]["duration"])
|
||||
|
||||
def get_image(self, idx) -> dict:
|
||||
image_data = dataset_prog.cap_list[
|
||||
idx] # [{'path': path, 'cap': cap}, ...]
|
||||
|
||||
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
|
||||
image = torch.from_numpy(np.array(image)) # [h, w, c]
|
||||
image = rearrange(image, "h w c -> c h w").unsqueeze(0) # [1 c h w]
|
||||
# for i in image:
|
||||
# h, w = i.shape[-2:]
|
||||
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
|
||||
|
||||
image = (self.transform_topcrop(image) if "human_images"
|
||||
in image_data["path"] else self.transform(image)
|
||||
) # [1 C H W] -> num_img [1 C H W]
|
||||
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
|
||||
|
||||
image = image.float() / 127.5 - 1.0
|
||||
|
||||
caps: list[str] = (image_data["cap"] if isinstance(
|
||||
image_data["cap"], list) else [image_data["cap"]])
|
||||
caps = [random.choice(caps)]
|
||||
text = caps
|
||||
input_ids, cond_mask = [], []
|
||||
single_text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
single_text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"] # 1, l
|
||||
cond_mask = text_tokens_and_mask["attention_mask"] # 1, l
|
||||
return dict(
|
||||
pixel_values=image,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=image_data["path"],
|
||||
)
|
||||
|
||||
def define_frame_index(self, cap_list) -> tuple[list[dict], list[int]]:
|
||||
new_cap_list = []
|
||||
sample_num_frames = []
|
||||
cnt_too_long = 0
|
||||
cnt_too_short = 0
|
||||
cnt_no_cap = 0
|
||||
cnt_no_resolution = 0
|
||||
cnt_resolution_mismatch = 0
|
||||
cnt_movie = 0
|
||||
cnt_img = 0
|
||||
for i in cap_list:
|
||||
path = i["path"]
|
||||
cap = i.get("cap", None)
|
||||
# ======no caption=====
|
||||
if cap is None:
|
||||
cnt_no_cap += 1
|
||||
continue
|
||||
if path.endswith(".mp4"):
|
||||
# ======no fps and duration=====
|
||||
duration = i.get("duration", None)
|
||||
fps = i.get("fps", None)
|
||||
if fps is None or duration is None:
|
||||
continue
|
||||
|
||||
# ======resolution mismatch=====
|
||||
resolution = i.get("resolution", None)
|
||||
if resolution is None:
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
else:
|
||||
if (resolution.get("height", None) is None
|
||||
or resolution.get("width", None) is None):
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
height, width = i["resolution"]["height"], i["resolution"][
|
||||
"width"]
|
||||
aspect = self.max_height / self.max_width
|
||||
hw_aspect_thr = 1.5
|
||||
is_pick = filter_resolution(
|
||||
height,
|
||||
width,
|
||||
max_h_div_w_ratio=hw_aspect_thr * aspect,
|
||||
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
|
||||
)
|
||||
if not is_pick:
|
||||
print("resolution mismatch")
|
||||
cnt_resolution_mismatch += 1
|
||||
continue
|
||||
|
||||
# if path == 'finetrainers/3dgs-dissolve/videos/1.mp4':
|
||||
# from IPython import embed; embed()
|
||||
i["num_frames"] = math.ceil(fps * duration)
|
||||
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
|
||||
if i["num_frames"] / fps > self.video_length_tolerance_range * (
|
||||
self.num_frames / self.train_fps * self.speed_factor
|
||||
): # too long video is not suitable for this training stage (self.num_frames)
|
||||
cnt_too_long += 1
|
||||
continue
|
||||
|
||||
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
|
||||
frame_interval = fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, i["num_frames"],
|
||||
frame_interval).astype(int)
|
||||
|
||||
# comment out it to enable dynamic frames training
|
||||
if (len(frame_indices) < self.num_frames
|
||||
and random.random() < self.drop_short_ratio):
|
||||
cnt_too_short += 1
|
||||
continue
|
||||
|
||||
# too long video will be temporal-crop randomly
|
||||
if len(frame_indices) > self.num_frames:
|
||||
begin_index, end_index = self.temporal_sample(
|
||||
len(frame_indices))
|
||||
frame_indices = frame_indices[begin_index:end_index]
|
||||
# frame_indices = frame_indices[:self.num_frames] # head crop
|
||||
i["sample_frame_index"] = frame_indices.tolist()
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = len(
|
||||
i["sample_frame_index"]
|
||||
) # will use in dataloader(group sampler)
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
elif path.endswith(".jpg"): # image
|
||||
cnt_img += 1
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = 1
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
else:
|
||||
raise NameError(
|
||||
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
|
||||
)
|
||||
# import ipdb;ipdb.set_trace()
|
||||
main_print(
|
||||
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
|
||||
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
|
||||
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
|
||||
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}"
|
||||
)
|
||||
return new_cap_list, sample_num_frames
|
||||
|
||||
def decord_read(self, path, frame_indices) -> torch.Tensor:
|
||||
decord_vr = self.v_decoder(path)
|
||||
video_data = decord_vr.get_batch(frame_indices).asnumpy()
|
||||
video_data = torch.from_numpy(video_data)
|
||||
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
|
||||
return video_data
|
||||
|
||||
def read_jsons(self, data) -> list[dict]:
|
||||
cap_lists = []
|
||||
with open(data) as f:
|
||||
folder_anno = [
|
||||
i.strip().split(",") for i in f.readlines()
|
||||
if len(i.strip()) > 0
|
||||
]
|
||||
print(folder_anno)
|
||||
for folder, anno in folder_anno:
|
||||
with open(anno) as f:
|
||||
sub_list = json.load(f)
|
||||
for i in range(len(sub_list)):
|
||||
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
|
||||
cap_lists += sub_list
|
||||
return cap_lists
|
||||
|
||||
def get_cap_list(self) -> list:
|
||||
cap_lists = self.read_jsons(self.data)[self.start_idx:]
|
||||
return cap_lists
|
||||
@@ -0,0 +1,220 @@
|
||||
import random
|
||||
from typing import Any, Dict, List, cast
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
|
||||
"""
|
||||
Pad or crop an embedding [L, D] to exactly padding_length tokens.
|
||||
Return:
|
||||
- [L, D] tensor in pinned CPU memory
|
||||
- [L] attention mask in pinned CPU memory
|
||||
"""
|
||||
L, D = t.shape
|
||||
if padding_length > L: # pad
|
||||
pad = torch.zeros(padding_length - L, D, dtype=t.dtype, device=t.device)
|
||||
return torch.cat([t, pad], 0), torch.cat(
|
||||
[torch.ones(L), torch.zeros(padding_length - L)], 0)
|
||||
else: # crop
|
||||
return t[:padding_length], torch.ones(padding_length)
|
||||
|
||||
|
||||
def get_torch_tensors_from_row_dict(row_dict, keys, cfg_rate) -> Dict[str, Any]:
|
||||
"""
|
||||
Get the latents and prompts from a row dictionary.
|
||||
"""
|
||||
return_dict = {}
|
||||
for key in keys:
|
||||
shape, bytes = None, None
|
||||
if isinstance(key, tuple):
|
||||
for k in key:
|
||||
try:
|
||||
shape = row_dict[f"{k}_shape"]
|
||||
bytes = row_dict[f"{k}_bytes"]
|
||||
except KeyError:
|
||||
continue
|
||||
key = key[0]
|
||||
if shape is None or bytes is None:
|
||||
raise ValueError(f"Key {key} not found in row_dict")
|
||||
else:
|
||||
shape = row_dict[f"{key}_shape"]
|
||||
bytes = row_dict[f"{key}_bytes"]
|
||||
|
||||
# TODO (peiyuan): read precision
|
||||
if key == 'text_embedding' and random.random() < cfg_rate:
|
||||
data = np.zeros((512, 4096), dtype=np.float32)
|
||||
else:
|
||||
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
|
||||
data = torch.from_numpy(data)
|
||||
if len(data.shape) == 3:
|
||||
B, L, D = data.shape
|
||||
assert B == 1, "Batch size must be 1"
|
||||
data = data.squeeze(0)
|
||||
return_dict[key] = data
|
||||
return return_dict
|
||||
|
||||
|
||||
def collate_latents_embs_masks(
|
||||
batch_to_process,
|
||||
text_padding_length,
|
||||
keys,
|
||||
cfg_rate=0.0
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, List[str]]:
|
||||
# Initialize tensors to hold padded embeddings and masks
|
||||
all_latents = []
|
||||
all_embs = []
|
||||
all_masks = []
|
||||
caption_text = []
|
||||
# Process each row individually
|
||||
for i, row in enumerate(batch_to_process):
|
||||
# Get tensors from row
|
||||
data = get_torch_tensors_from_row_dict(row, keys, cfg_rate)
|
||||
latents, emb = data["vae_latent"], data["text_embedding"]
|
||||
|
||||
padded_emb, mask = pad(emb, text_padding_length)
|
||||
# Store in batch tensors
|
||||
all_latents.append(latents)
|
||||
all_embs.append(padded_emb)
|
||||
all_masks.append(mask)
|
||||
# TODO(py): remove this once we fix preprocess
|
||||
try:
|
||||
caption_text.append(row["prompt"])
|
||||
except KeyError:
|
||||
caption_text.append(row["caption"])
|
||||
|
||||
# Pin memory for faster transfer to GPU
|
||||
all_latents = torch.stack(all_latents)
|
||||
all_embs = torch.stack(all_embs)
|
||||
all_masks = torch.stack(all_masks)
|
||||
|
||||
return all_latents, all_embs, all_masks, caption_text
|
||||
|
||||
|
||||
def collate_rows_from_parquet_schema(rows,
|
||||
parquet_schema,
|
||||
text_padding_length,
|
||||
cfg_rate=0.0) -> Dict[str, Any]:
|
||||
"""
|
||||
Collate rows from parquet files based on the provided schema.
|
||||
Dynamically processes tensor fields based on schema and returns batched data.
|
||||
|
||||
Args:
|
||||
rows: List of row dictionaries from parquet files
|
||||
parquet_schema: PyArrow schema defining the structure of the data
|
||||
|
||||
Returns:
|
||||
Dict containing batched tensors and metadata
|
||||
"""
|
||||
if not rows:
|
||||
return cast(Dict[str, Any], {})
|
||||
|
||||
# Initialize containers for different data types
|
||||
batch_data: Dict[str, Any] = {}
|
||||
|
||||
# Get tensor and metadata field names from schema (fields ending with '_bytes')
|
||||
tensor_fields = []
|
||||
metadata_fields = []
|
||||
for field in parquet_schema.names:
|
||||
if field.endswith('_bytes'):
|
||||
shape_field = field.replace('_bytes', '_shape')
|
||||
dtype_field = field.replace('_bytes', '_dtype')
|
||||
tensor_name = field.replace('_bytes', '')
|
||||
tensor_fields.append(tensor_name)
|
||||
assert shape_field in parquet_schema.names, f"Shape field {shape_field} not found in schema for field {field}. Currently we only support *_bytes fields for tensors."
|
||||
assert dtype_field in parquet_schema.names, f"Dtype field {dtype_field} not found in schema for field {field}. Currently we only support *_bytes fields for tensors."
|
||||
elif not field.endswith('_shape') and not field.endswith('_dtype'):
|
||||
# Only add actual metadata fields, not the shape/dtype helper fields
|
||||
metadata_fields.append(field)
|
||||
|
||||
# Process each tensor field
|
||||
for tensor_name in tensor_fields:
|
||||
tensor_list = []
|
||||
|
||||
for row in rows:
|
||||
# Get tensor data from row using the existing helper function pattern
|
||||
shape_key = f"{tensor_name}_shape"
|
||||
bytes_key = f"{tensor_name}_bytes"
|
||||
|
||||
if shape_key in row and bytes_key in row:
|
||||
shape = row[shape_key]
|
||||
bytes_data = row[bytes_key]
|
||||
|
||||
if len(bytes_data) == 0:
|
||||
tensor = torch.zeros(0, dtype=torch.bfloat16)
|
||||
else:
|
||||
# Convert bytes to tensor using float32 as default
|
||||
if tensor_name == 'text_embedding' and random.random(
|
||||
) < cfg_rate:
|
||||
data = np.zeros((512, 4096), dtype=np.float32)
|
||||
else:
|
||||
data = np.frombuffer(
|
||||
bytes_data, dtype=np.float32).reshape(shape).copy()
|
||||
tensor = torch.from_numpy(data)
|
||||
# if len(data.shape) == 3:
|
||||
# B, L, D = tensor.shape
|
||||
# assert B == 1, "Batch size must be 1"
|
||||
# tensor = tensor.squeeze(0)
|
||||
|
||||
tensor_list.append(tensor)
|
||||
else:
|
||||
# Handle missing tensor data
|
||||
tensor_list.append(torch.zeros(0, dtype=torch.bfloat16))
|
||||
|
||||
# Stack tensors with special handling for text embeddings
|
||||
if tensor_name == 'text_embedding':
|
||||
# Handle text embeddings with padding
|
||||
padded_tensors = []
|
||||
attention_masks = []
|
||||
|
||||
for tensor in tensor_list:
|
||||
if tensor.numel() > 0:
|
||||
padded_tensor, mask = pad(tensor, text_padding_length)
|
||||
padded_tensors.append(padded_tensor)
|
||||
attention_masks.append(mask)
|
||||
else:
|
||||
# Handle empty embeddings - assume default embedding dimension
|
||||
padded_tensors.append(
|
||||
torch.zeros(text_padding_length,
|
||||
768,
|
||||
dtype=torch.bfloat16))
|
||||
attention_masks.append(torch.zeros(text_padding_length))
|
||||
|
||||
batch_data[tensor_name] = torch.stack(padded_tensors)
|
||||
batch_data['text_attention_mask'] = torch.stack(attention_masks)
|
||||
else:
|
||||
# Stack all tensors to preserve batch consistency
|
||||
# Don't filter out None or empty tensors as this breaks batch sizing
|
||||
try:
|
||||
batch_data[tensor_name] = torch.stack(tensor_list)
|
||||
except ValueError as e:
|
||||
shapes = [
|
||||
t.shape
|
||||
if t is not None and hasattr(t, 'shape') else 'None/Invalid'
|
||||
for t in tensor_list
|
||||
]
|
||||
raise ValueError(
|
||||
f"Failed to stack tensors for field '{tensor_name}'. "
|
||||
f"Tensor shapes: {shapes}. "
|
||||
f"All tensors in a batch must have compatible shapes. "
|
||||
f"Original error: {e}") from e
|
||||
|
||||
# Process metadata fields into info_list
|
||||
info_list = []
|
||||
for row in rows:
|
||||
info = {}
|
||||
for field in metadata_fields:
|
||||
info[field] = row.get(field, "")
|
||||
|
||||
# Add prompt field for backward compatibility
|
||||
info["prompt"] = info.get("caption", "")
|
||||
info_list.append(info)
|
||||
|
||||
batch_data['info_list'] = info_list
|
||||
|
||||
# Add caption_text for backward compatibility
|
||||
if info_list and 'caption' in info_list[0]:
|
||||
batch_data['caption_text'] = [info['caption'] for info in info_list]
|
||||
|
||||
return batch_data
|
||||
@@ -0,0 +1,103 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from: https://github.com/a-r-r-o-w/finetrainers/blob/main/finetrainers/data/dataset.py
|
||||
import os
|
||||
import pathlib
|
||||
|
||||
import datasets
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vision_utils import load_image, load_video
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ValidationDataset(torch.utils.data.IterableDataset):
|
||||
|
||||
def __init__(self, filename: str):
|
||||
super().__init__()
|
||||
|
||||
self.filename = pathlib.Path(filename)
|
||||
# get directory of filename
|
||||
self.dir = os.path.abspath(self.filename.parent)
|
||||
|
||||
if not self.filename.exists():
|
||||
raise FileNotFoundError(
|
||||
f"File {self.filename.as_posix()} does not exist")
|
||||
|
||||
if self.filename.suffix == ".csv":
|
||||
data = datasets.load_dataset("csv",
|
||||
data_files=self.filename.as_posix(),
|
||||
split="train")
|
||||
elif self.filename.suffix == ".json":
|
||||
data = datasets.load_dataset("json",
|
||||
data_files=self.filename.as_posix(),
|
||||
split="train",
|
||||
field="data")
|
||||
elif self.filename.suffix == ".parquet":
|
||||
data = datasets.load_dataset("parquet",
|
||||
data_files=self.filename.as_posix(),
|
||||
split="train")
|
||||
elif self.filename.suffix == ".arrow":
|
||||
data = datasets.load_dataset("arrow",
|
||||
data_files=self.filename.as_posix(),
|
||||
split="train")
|
||||
else:
|
||||
_SUPPORTED_FILE_FORMATS = [".csv", ".json", ".parquet", ".arrow"]
|
||||
raise ValueError(
|
||||
f"Unsupported file format {self.filename.suffix} for validation dataset. Supported formats are: {_SUPPORTED_FILE_FORMATS}"
|
||||
)
|
||||
|
||||
self._data = data.to_iterable_dataset()
|
||||
|
||||
def __iter__(self):
|
||||
for sample in self._data:
|
||||
# For consistency reasons, we mandate that "caption" is always present in the validation dataset.
|
||||
# However, since the model specifications use "prompt", we create an alias here.
|
||||
sample["prompt"] = sample["caption"]
|
||||
|
||||
# Load image or video if the path is provided
|
||||
# TODO(aryan): need to handle custom columns here for control conditions
|
||||
sample["image"] = None
|
||||
sample["video"] = None
|
||||
|
||||
if sample.get("image_path", None) is not None:
|
||||
image_path = sample["image_path"]
|
||||
image_path = os.path.join(self.dir, image_path)
|
||||
if not pathlib.Path(image_path).is_file(
|
||||
) and not image_path.startswith("http"):
|
||||
logger.warning("Image file %s does not exist.", image_path)
|
||||
else:
|
||||
sample["image"] = load_image(image_path)
|
||||
|
||||
if sample.get("video_path", None) is not None:
|
||||
video_path = sample["video_path"]
|
||||
video_path = os.path.join(self.dir, video_path)
|
||||
if not pathlib.Path(video_path).is_file(
|
||||
) and not video_path.startswith("http"):
|
||||
logger.warning("Video file %s does not exist.", video_path)
|
||||
else:
|
||||
sample["video"] = load_video(video_path)
|
||||
|
||||
if sample.get("control_image_path", None) is not None:
|
||||
control_image_path = sample["control_image_path"]
|
||||
control_image_path = os.path.join(self.dir, control_image_path)
|
||||
if not pathlib.Path(control_image_path).is_file(
|
||||
) and not control_image_path.startswith("http"):
|
||||
logger.warning("Control Image file %s does not exist.",
|
||||
control_image_path)
|
||||
else:
|
||||
sample["control_image"] = load_image(control_image_path)
|
||||
|
||||
if sample.get("control_video_path", None) is not None:
|
||||
control_video_path = sample["control_video_path"]
|
||||
control_video_path = os.path.join(self.dir, control_video_path)
|
||||
if not pathlib.Path(control_video_path).is_file(
|
||||
) and not control_video_path.startswith("http"):
|
||||
logger.warning("Control Video file %s does not exist.",
|
||||
control_video_path)
|
||||
else:
|
||||
sample["control_video"] = load_video(control_video_path)
|
||||
|
||||
sample = {k: v for k, v in sample.items() if v is not None}
|
||||
yield sample
|
||||
@@ -3,10 +3,10 @@
|
||||
from fastvideo.v1.distributed.communication_op import *
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_dp_group, get_dp_rank, get_dp_world_size,
|
||||
get_sp_group, get_sp_parallel_rank, get_sp_world_size, get_torch_device,
|
||||
get_tp_group, get_tp_rank, get_tp_world_size, get_world_group,
|
||||
get_world_rank, get_world_size, init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
get_local_torch_device, get_sp_group, get_sp_parallel_rank,
|
||||
get_sp_world_size, get_tp_group, get_tp_rank, get_tp_world_size,
|
||||
get_world_group, get_world_rank, get_world_size,
|
||||
init_distributed_environment, initialize_model_parallel,
|
||||
maybe_init_distributed_environment_and_model_parallel,
|
||||
model_parallel_is_initialized)
|
||||
from fastvideo.v1.distributed.utils import *
|
||||
@@ -40,5 +40,5 @@ __all__ = [
|
||||
"get_tp_world_size",
|
||||
|
||||
# Get torch device
|
||||
"get_torch_device",
|
||||
"get_local_torch_device",
|
||||
]
|
||||
|
||||
@@ -36,6 +36,7 @@ from unittest.mock import patch
|
||||
|
||||
import torch
|
||||
import torch.distributed
|
||||
import torch.distributed as dist
|
||||
from torch.distributed import Backend, ProcessGroup, ReduceOp
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
@@ -692,6 +693,7 @@ class GroupCoordinator:
|
||||
|
||||
|
||||
_WORLD: Optional[GroupCoordinator] = None
|
||||
_NODE: Optional[GroupCoordinator] = None
|
||||
|
||||
|
||||
def get_world_group() -> GroupCoordinator:
|
||||
@@ -699,6 +701,11 @@ def get_world_group() -> GroupCoordinator:
|
||||
return _WORLD
|
||||
|
||||
|
||||
def get_node_group() -> GroupCoordinator:
|
||||
assert _NODE is not None, ("node group is not initialized")
|
||||
return _NODE
|
||||
|
||||
|
||||
def init_world_group(ranks: List[int], local_rank: int,
|
||||
backend: str) -> GroupCoordinator:
|
||||
return GroupCoordinator(
|
||||
@@ -710,6 +717,18 @@ def init_world_group(ranks: List[int], local_rank: int,
|
||||
)
|
||||
|
||||
|
||||
def init_node_group(local_rank: int, backend: str):
|
||||
cpu_group = get_world_group().cpu_group
|
||||
node_ranks = same_node_ranks(cpu_group)
|
||||
node_size = len(node_ranks)
|
||||
all_node_ranks = [
|
||||
list(range(i * node_size, (i + 1) * node_size))
|
||||
for i in range(dist.get_world_size() // node_size)
|
||||
]
|
||||
global _NODE
|
||||
_NODE = init_model_parallel_group(all_node_ranks, local_rank, backend)
|
||||
|
||||
|
||||
def init_model_parallel_group(
|
||||
group_ranks: List[List[int]],
|
||||
local_rank: int,
|
||||
@@ -782,6 +801,8 @@ def init_distributed_environment(
|
||||
else:
|
||||
assert _WORLD.world_size == torch.distributed.get_world_size(), (
|
||||
"world group already initialized with a different world size")
|
||||
# Init a group for each node
|
||||
init_node_group(local_rank, backend)
|
||||
|
||||
|
||||
_SP: Optional[GroupCoordinator] = None
|
||||
@@ -904,7 +925,7 @@ def get_dp_rank() -> int:
|
||||
return get_dp_group().rank_in_group
|
||||
|
||||
|
||||
def get_torch_device() -> torch.device:
|
||||
def get_local_torch_device() -> torch.device:
|
||||
"""Return the torch device for the current rank."""
|
||||
return torch.device(f"cuda:{envs.LOCAL_RANK}")
|
||||
|
||||
@@ -1021,17 +1042,22 @@ def cleanup_dist_env_and_memory(shutdown_ray: bool = False):
|
||||
"torch._C._host_emptyCache() only available in Pytorch >=2.5")
|
||||
|
||||
|
||||
def in_the_same_node_as(pg: Union[ProcessGroup, StatelessProcessGroup],
|
||||
source_rank: int = 0) -> List[bool]:
|
||||
def same_node_ranks(pg: Union[ProcessGroup, StatelessProcessGroup],
|
||||
source_rank: int = 0) -> List[int]:
|
||||
"""
|
||||
This is a collective operation that returns if each rank is in the same node
|
||||
This is a collective operation that returns ranks that are in the same node
|
||||
as the source rank. It tests if processes are attached to the same
|
||||
memory system (shared access to shared memory).
|
||||
Args:
|
||||
pg: the global process group to test
|
||||
source_rank: the rank to test against
|
||||
Returns:
|
||||
A list of ranks that are in the same node as the source rank.
|
||||
"""
|
||||
if isinstance(pg, ProcessGroup):
|
||||
assert torch.distributed.get_backend(
|
||||
pg) != torch.distributed.Backend.NCCL, (
|
||||
"in_the_same_node_as should be tested with a non-NCCL group.")
|
||||
"same_node_ranks should be tested with a non-NCCL group.")
|
||||
# local rank inside the group
|
||||
rank = torch.distributed.get_rank(group=pg)
|
||||
world_size = torch.distributed.get_world_size(group=pg)
|
||||
@@ -1103,7 +1129,7 @@ def in_the_same_node_as(pg: Union[ProcessGroup, StatelessProcessGroup],
|
||||
rank_data = pg.broadcast_obj(is_in_the_same_node, src=i)
|
||||
aggregated_data += rank_data
|
||||
|
||||
return [x == 1 for x in aggregated_data.tolist()]
|
||||
return [i for i, x in enumerate(aggregated_data.tolist()) if x == 1]
|
||||
|
||||
|
||||
def initialize_tensor_parallel_group(
|
||||
|
||||
@@ -4,9 +4,9 @@
|
||||
import argparse
|
||||
import dataclasses
|
||||
import os
|
||||
from typing import Any, Dict, List, Optional, cast
|
||||
from typing import List, cast
|
||||
|
||||
from fastvideo import PipelineConfig, VideoGenerator
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.v1.entrypoints.cli.utils import RaiseNotImplementedAction
|
||||
@@ -37,8 +37,6 @@ class GenerateSubcommand(CLISubcommand):
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
excluded_args = ['subparser', 'config', 'dispatch_function']
|
||||
|
||||
FastVideoArgs.from_cli_args(args)
|
||||
|
||||
provided_args = {}
|
||||
for k, v in vars(args).items():
|
||||
if (k not in excluded_args and v is not None
|
||||
@@ -66,27 +64,19 @@ class GenerateSubcommand(CLISubcommand):
|
||||
|
||||
init_args = {
|
||||
k: v
|
||||
for k, v in merged_args.items() if k in self.init_arg_names
|
||||
for k, v in merged_args.items()
|
||||
if k not in self.generation_arg_names
|
||||
}
|
||||
generation_args = {
|
||||
k: v
|
||||
for k, v in merged_args.items() if k in self.generation_arg_names
|
||||
}
|
||||
|
||||
pipeline_config = PipelineConfig.from_pretrained(
|
||||
merged_args['model_path'])
|
||||
|
||||
update_config_from_args(pipeline_config.dit_config, merged_args,
|
||||
"dit_config")
|
||||
update_config_from_args(pipeline_config.vae_config, merged_args,
|
||||
"vae_config")
|
||||
update_config_from_args(pipeline_config, merged_args)
|
||||
|
||||
model_path = init_args.pop('model_path')
|
||||
prompt = generation_args.pop('prompt')
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=model_path, **init_args, pipeline_config=pipeline_config)
|
||||
generator = VideoGenerator.from_pretrained(model_path=model_path,
|
||||
**init_args)
|
||||
|
||||
generator.generate_video(prompt=prompt, **generation_args)
|
||||
|
||||
@@ -132,34 +122,3 @@ class GenerateSubcommand(CLISubcommand):
|
||||
|
||||
def cmd_init() -> List[CLISubcommand]:
|
||||
return [GenerateSubcommand()]
|
||||
|
||||
|
||||
def update_config_from_args(config: Any,
|
||||
args_dict: Dict[str, Any],
|
||||
prefix: Optional[str] = None) -> None:
|
||||
"""
|
||||
Update configuration object from arguments dictionary.
|
||||
|
||||
Args:
|
||||
config: The configuration object to update
|
||||
args_dict: Dictionary containing arguments
|
||||
prefix: Prefix for the configuration parameters in the args_dict.
|
||||
If None, assumes direct attribute mapping without prefix.
|
||||
"""
|
||||
# Handle top-level attributes (no prefix)
|
||||
if prefix is None:
|
||||
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)
|
||||
return
|
||||
|
||||
# Handle nested attributes with prefix
|
||||
prefix_with_dot = f"{prefix}."
|
||||
for key, value in args_dict.items():
|
||||
if key.startswith(prefix_with_dot) and value is not None:
|
||||
attr_name = key[len(prefix_with_dot):]
|
||||
if hasattr(config, attr_name):
|
||||
setattr(config, attr_name, value)
|
||||
|
||||
@@ -18,8 +18,6 @@ import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.v1.configs.pipelines import (PipelineConfig,
|
||||
get_pipeline_config_cls_for_name)
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -55,9 +53,6 @@ class VideoGenerator:
|
||||
model_path: str,
|
||||
device: Optional[str] = None,
|
||||
torch_dtype: Optional[torch.dtype] = None,
|
||||
pipeline_config: Optional[
|
||||
Union[str
|
||||
| PipelineConfig]] = None,
|
||||
**kwargs) -> "VideoGenerator":
|
||||
"""
|
||||
Create a video generator from a pretrained model.
|
||||
@@ -66,35 +61,17 @@ class VideoGenerator:
|
||||
model_path: Path or identifier for the pretrained model
|
||||
device: Device to load the model on (e.g., "cuda", "cuda:0", "cpu")
|
||||
torch_dtype: Data type for model weights (e.g., torch.float16)
|
||||
**kwargs: Additional arguments to customize model loading
|
||||
pipeline_config: Pipeline config to use for inference
|
||||
**kwargs: Additional arguments to customize model loading, set any FastVideoArgs or PipelineConfig attributes here.
|
||||
|
||||
Returns:
|
||||
The created video generator
|
||||
|
||||
Priority level: Default pipeline config < User's pipeline config < User's kwargs
|
||||
"""
|
||||
config = None
|
||||
# 1. If users provide a pipeline config, it will override the default pipeline config
|
||||
if isinstance(pipeline_config, PipelineConfig):
|
||||
config = pipeline_config
|
||||
else:
|
||||
config_cls = get_pipeline_config_cls_for_name(model_path)
|
||||
if config_cls is not None:
|
||||
config = config_cls()
|
||||
if isinstance(pipeline_config, str):
|
||||
config.load_from_json(pipeline_config)
|
||||
|
||||
# 2. If users also provide some kwargs, it will override the pipeline config.
|
||||
# The user kwargs shouldn't contain model config parameters!
|
||||
if config is None:
|
||||
logger.warning("No config found for model %s, using default config",
|
||||
model_path)
|
||||
config_args = kwargs
|
||||
else:
|
||||
config_args = shallow_asdict(config)
|
||||
config_args.update(kwargs)
|
||||
|
||||
fastvideo_args = FastVideoArgs(model_path=model_path, **config_args)
|
||||
# If users also provide some kwargs, it will override the FastVideoArgs and PipelineConfig.
|
||||
kwargs['model_path'] = model_path
|
||||
fastvideo_args = FastVideoArgs.from_kwargs(kwargs)
|
||||
|
||||
return cls.from_fastvideo_args(fastvideo_args)
|
||||
|
||||
@@ -150,16 +127,17 @@ class VideoGenerator:
|
||||
"""
|
||||
# Create a copy of inference args to avoid modifying the original
|
||||
fastvideo_args = self.fastvideo_args
|
||||
pipeline_config = fastvideo_args.pipeline_config
|
||||
|
||||
# Validate inputs
|
||||
if not isinstance(prompt, str):
|
||||
raise TypeError(
|
||||
f"`prompt` must be a string, but got {type(prompt)}")
|
||||
prompt = prompt.strip()
|
||||
|
||||
if sampling_param is None:
|
||||
sampling_param = SamplingParam.from_pretrained(
|
||||
fastvideo_args.model_path)
|
||||
|
||||
kwargs["prompt"] = prompt
|
||||
sampling_param.update(kwargs)
|
||||
|
||||
@@ -176,10 +154,10 @@ class VideoGenerator:
|
||||
f"height={sampling_param.height}, width={sampling_param.width}, "
|
||||
f"num_frames={sampling_param.num_frames}")
|
||||
|
||||
temporal_scale_factor = fastvideo_args.vae_config.arch_config.temporal_compression_ratio
|
||||
temporal_scale_factor = pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
num_frames = sampling_param.num_frames
|
||||
num_gpus = fastvideo_args.num_gpus
|
||||
use_temporal_scaling_frames = fastvideo_args.vae_config.use_temporal_scaling_frames
|
||||
use_temporal_scaling_frames = pipeline_config.vae_config.use_temporal_scaling_frames
|
||||
|
||||
# Adjust number of frames based on number of GPUs
|
||||
if use_temporal_scaling_frames:
|
||||
@@ -238,18 +216,18 @@ class VideoGenerator:
|
||||
num_videos_per_prompt: {sampling_param.num_videos_per_prompt}
|
||||
guidance_scale: {sampling_param.guidance_scale}
|
||||
n_tokens: {n_tokens}
|
||||
flow_shift: {fastvideo_args.flow_shift}
|
||||
embedded_guidance_scale: {fastvideo_args.embedded_cfg_scale}
|
||||
flow_shift: {fastvideo_args.pipeline_config.flow_shift}
|
||||
embedded_guidance_scale: {fastvideo_args.pipeline_config.embedded_cfg_scale}
|
||||
save_video: {sampling_param.save_video}
|
||||
output_path: {sampling_param.output_path}
|
||||
""" # type: ignore[attr-defined]
|
||||
logger.info(debug_str)
|
||||
|
||||
# Prepare batch
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
eta=0.0,
|
||||
n_tokens=n_tokens,
|
||||
VSA_sparsity=fastvideo_args.VSA_sparsity,
|
||||
extra={},
|
||||
)
|
||||
|
||||
|
||||
+125
-239
@@ -6,26 +6,32 @@ import argparse
|
||||
import dataclasses
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import field
|
||||
from typing import Any, Callable, List, Optional, Tuple
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig, STA_Mode
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser, StoreBoolean
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def preprocess_text(prompt: str) -> str:
|
||||
return prompt
|
||||
def clean_cli_args(args: argparse.Namespace) -> Dict[str, Any]:
|
||||
"""
|
||||
Clean the arguments by removing the ones that not explicitly provided by the user.
|
||||
"""
|
||||
provided_args = {}
|
||||
for k, v in vars(args).items():
|
||||
if (v is not None and hasattr(args, '_provided')
|
||||
and k in args._provided):
|
||||
provided_args[k] = v
|
||||
|
||||
|
||||
def postprocess_text(output: Any) -> Any:
|
||||
raise NotImplementedError
|
||||
return provided_args
|
||||
|
||||
|
||||
# args for fastvideo framework
|
||||
@dataclasses.dataclass
|
||||
class FastVideoArgs:
|
||||
# Model and path configuration
|
||||
# Model and path configuration (for convenience)
|
||||
model_path: str
|
||||
|
||||
# Cache strategy
|
||||
@@ -44,70 +50,34 @@ class FastVideoArgs:
|
||||
num_gpus: int = 1
|
||||
tp_size: int = -1
|
||||
sp_size: int = -1
|
||||
dp_size: int = 1
|
||||
dp_shards: int = -1
|
||||
hsdp_replicate_dim: int = 1
|
||||
hsdp_shard_dim: int = -1
|
||||
dist_timeout: Optional[int] = None # timeout for torch.distributed
|
||||
|
||||
# Video generation parameters
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: Optional[float] = None
|
||||
pipeline_config: PipelineConfig = field(default_factory=PipelineConfig)
|
||||
|
||||
output_type: str = "pil"
|
||||
|
||||
# DiT configuration
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
precision: str = "bf16"
|
||||
use_cpu_offload: bool = True
|
||||
use_cpu_offload: bool = True # For DiT
|
||||
use_fsdp_inference: bool = True
|
||||
text_encoder_offload: bool = True
|
||||
pin_cpu_memory: bool = True
|
||||
|
||||
# VAE configuration
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True # Might change in between forward passes
|
||||
vae_sp: bool = False # Might change in between forward passes
|
||||
# vae_scale_factor: Optional[int] = None # Deprecated
|
||||
vae_config: VAEConfig = field(default_factory=VAEConfig)
|
||||
|
||||
# Image encoder configuration
|
||||
image_encoder_precision: str = "fp32"
|
||||
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
|
||||
|
||||
# Text encoder configuration
|
||||
DEFAULT_TEXT_ENCODER_PRECISIONS = (
|
||||
"fp16",
|
||||
# "fp16",
|
||||
)
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS)
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (EncoderConfig(), ))
|
||||
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (preprocess_text, ))
|
||||
postprocess_text_funcs: Tuple[Callable[[Any], Any], ...] = field(
|
||||
default_factory=lambda: (postprocess_text, ))
|
||||
|
||||
# STA parameters
|
||||
STA_mode: Optional[str] = None
|
||||
skip_time_steps: int = 15
|
||||
# LoRA parameters
|
||||
lora_path: Optional[str] = None
|
||||
lora_nickname: Optional[
|
||||
str] = "default" # for swapping adapters in the pipeline
|
||||
lora_target_names: Optional[List[
|
||||
str]] = None # can restrict list of layers to adapt, e.g. ["q_proj"]
|
||||
|
||||
# STA parameters
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
STA_mode: STA_Mode = STA_Mode.STA_INFERENCE
|
||||
skip_time_steps: int = 15
|
||||
|
||||
# Compilation
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
disable_autocast: bool = False
|
||||
|
||||
# StepVideo specific parameters
|
||||
pos_magic: Optional[str] = None
|
||||
neg_magic: Optional[str] = None
|
||||
timesteps_scale: Optional[bool] = None
|
||||
# VSA parameters
|
||||
VSA_sparsity: float = 0.0 # inference/validation sparsity
|
||||
|
||||
# Logging
|
||||
log_level: str = "info"
|
||||
# Stage verification
|
||||
enable_stage_verification: bool = True
|
||||
|
||||
@property
|
||||
def training_mode(self) -> bool:
|
||||
@@ -125,11 +95,6 @@ class FastVideoArgs:
|
||||
help=
|
||||
"The path of the model weights. This can be a local folder or a Hugging Face repo ID.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dit-weight",
|
||||
type=str,
|
||||
help="Path to the DiT model weights",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model-dir",
|
||||
type=str,
|
||||
@@ -175,31 +140,27 @@ class FastVideoArgs:
|
||||
help="The number of GPUs to use.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tensor-parallel-size",
|
||||
"--tp-size",
|
||||
type=int,
|
||||
default=FastVideoArgs.tp_size,
|
||||
help="The tensor parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sequence-parallel-size",
|
||||
"--sp-size",
|
||||
type=int,
|
||||
default=FastVideoArgs.sp_size,
|
||||
help="The sequence parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--data-parallel-size",
|
||||
"--dp-size",
|
||||
"--hsdp-replicate-dim",
|
||||
type=int,
|
||||
default=FastVideoArgs.dp_size,
|
||||
default=FastVideoArgs.hsdp_replicate_dim,
|
||||
help="The data parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--data-parallel-shards",
|
||||
"--dp-shards",
|
||||
"--hsdp-shard-dim",
|
||||
type=int,
|
||||
default=FastVideoArgs.dp_shards,
|
||||
default=FastVideoArgs.hsdp_shard_dim,
|
||||
help="The data parallelism shards.",
|
||||
)
|
||||
parser.add_argument(
|
||||
@@ -209,19 +170,7 @@ class FastVideoArgs:
|
||||
help="Set timeout for torch.distributed initialization.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--embedded-cfg-scale",
|
||||
type=float,
|
||||
default=FastVideoArgs.embedded_cfg_scale,
|
||||
help="Embedded CFG scale",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flow-shift",
|
||||
"--shift",
|
||||
type=float,
|
||||
default=FastVideoArgs.flow_shift,
|
||||
help="Flow shift parameter",
|
||||
)
|
||||
# Output type
|
||||
parser.add_argument(
|
||||
"--output-type",
|
||||
type=str,
|
||||
@@ -230,62 +179,14 @@ class FastVideoArgs:
|
||||
help="Output type for the generated video",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--precision",
|
||||
type=str,
|
||||
default=FastVideoArgs.precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for the model",
|
||||
)
|
||||
|
||||
# VAE configuration
|
||||
parser.add_argument(
|
||||
"--vae-precision",
|
||||
type=str,
|
||||
default=FastVideoArgs.vae_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for VAE",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-tiling",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.vae_tiling,
|
||||
help="Enable VAE tiling",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-sp",
|
||||
action=StoreBoolean,
|
||||
help="Enable VAE spatial parallelism",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--text-encoder-precisions",
|
||||
nargs="+",
|
||||
type=str,
|
||||
default=FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for each text encoder",
|
||||
)
|
||||
|
||||
# Image encoder config
|
||||
parser.add_argument(
|
||||
"--image-encoder-precision",
|
||||
type=str,
|
||||
default=FastVideoArgs.image_encoder_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for image encoder",
|
||||
)
|
||||
|
||||
# STA parameters
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
parser.add_argument(
|
||||
"--STA-mode",
|
||||
type=str,
|
||||
default=FastVideoArgs.STA_mode,
|
||||
choices=[
|
||||
"STA_inference", "STA_searching", "STA_tuning",
|
||||
"STA_tuning_cfg", None
|
||||
],
|
||||
help="STA mode",
|
||||
default=FastVideoArgs.STA_mode.value,
|
||||
choices=[mode.value for mode in STA_Mode],
|
||||
help=
|
||||
"STA mode contains STA_inference, STA_searching, STA_tuning, STA_tuning_cfg, None",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--skip-time-steps",
|
||||
@@ -309,7 +210,7 @@ class FastVideoArgs:
|
||||
"--use-cpu-offload",
|
||||
action=StoreBoolean,
|
||||
help=
|
||||
"Use CPU offload for model inference. Enable if run out of memory with FSDP.",
|
||||
"Use CPU offload for DiT inference. Enable if run out of memory with FSDP.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-fsdp-inference",
|
||||
@@ -317,7 +218,19 @@ class FastVideoArgs:
|
||||
help=
|
||||
"Use FSDP for inference by sharding the model weights. Latency is very low due to prefetch--enable if run out of memory.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--text-encoder-cpu-offload",
|
||||
action=StoreBoolean,
|
||||
help=
|
||||
"Use CPU offload for text encoder. Enable if run out of memory.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--pin-cpu-memory",
|
||||
action=StoreBoolean,
|
||||
help=
|
||||
"Pin memory for CPU offload. Only added as a temp workaround if it throws \"CUDA error: invalid argument\". "
|
||||
"Should be enabled in almost all cases",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--disable-autocast",
|
||||
action=StoreBoolean,
|
||||
@@ -325,91 +238,69 @@ class FastVideoArgs:
|
||||
"Disable autocast for denoising loop and vae decoding in pipeline sampling",
|
||||
)
|
||||
|
||||
# VSA parameters
|
||||
parser.add_argument(
|
||||
"--pos_magic",
|
||||
type=str,
|
||||
default=FastVideoArgs.pos_magic,
|
||||
help="Positive magic prompt for sampling",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--neg_magic",
|
||||
type=str,
|
||||
default=FastVideoArgs.neg_magic,
|
||||
help="Negative magic prompt for sampling",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--timesteps_scale",
|
||||
type=bool,
|
||||
default=FastVideoArgs.timesteps_scale,
|
||||
help="Bool for applying scheduler scale in set_timesteps",
|
||||
"--VSA-sparsity",
|
||||
type=float,
|
||||
default=FastVideoArgs.VSA_sparsity,
|
||||
help="Validation sparsity for VSA",
|
||||
)
|
||||
|
||||
# Logging
|
||||
# Stage verification
|
||||
parser.add_argument(
|
||||
"--log-level",
|
||||
type=str,
|
||||
default=FastVideoArgs.log_level,
|
||||
help="The logging level of all loggers.",
|
||||
"--enable-stage-verification",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.enable_stage_verification,
|
||||
help="Enable input/output verification for pipeline stages",
|
||||
)
|
||||
|
||||
# Add VAE configuration arguments
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEConfig
|
||||
VAEConfig.add_cli_args(parser)
|
||||
|
||||
# Add DiT configuration arguments
|
||||
from fastvideo.v1.configs.models.dits.base import DiTConfig
|
||||
DiTConfig.add_cli_args(parser)
|
||||
# Add pipeline configuration arguments
|
||||
PipelineConfig.add_cli_args(parser)
|
||||
|
||||
return parser
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "FastVideoArgs":
|
||||
args.tp_size = args.tensor_parallel_size
|
||||
args.sp_size = args.sequence_parallel_size
|
||||
args.flow_shift = getattr(args, "shift", args.flow_shift)
|
||||
|
||||
provided_args = clean_cli_args(args)
|
||||
# Get all fields from the dataclass
|
||||
attrs = [attr.name for attr in dataclasses.fields(cls)]
|
||||
|
||||
# Create a dictionary of attribute values, with defaults for missing attributes
|
||||
kwargs = {}
|
||||
for attr in attrs:
|
||||
# Handle renamed attributes or those with multiple CLI names
|
||||
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
|
||||
kwargs[attr] = args.tensor_parallel_size
|
||||
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
|
||||
kwargs[attr] = args.sequence_parallel_size
|
||||
elif attr == 'dp_size' and hasattr(args, 'data_parallel_size'):
|
||||
kwargs[attr] = args.data_parallel_size
|
||||
elif attr == 'dp_shards' and hasattr(args, 'data_parallel_shards'):
|
||||
kwargs[attr] = args.data_parallel_shards
|
||||
elif attr == 'flow_shift' and hasattr(args, 'shift'):
|
||||
kwargs[attr] = args.shift
|
||||
if attr == 'pipeline_config':
|
||||
pipeline_config = PipelineConfig.from_kwargs(provided_args)
|
||||
kwargs[attr] = pipeline_config
|
||||
# Use getattr with default value from the dataclass for potentially missing attributes
|
||||
else:
|
||||
default_value = getattr(cls, attr, None)
|
||||
value = getattr(args, attr, default_value)
|
||||
if value is not None:
|
||||
kwargs[attr] = value
|
||||
kwargs[attr] = value # type: ignore
|
||||
|
||||
return cls(**kwargs) # type: ignore
|
||||
|
||||
@classmethod
|
||||
def from_kwargs(cls, kwargs: Dict[str, Any]) -> "FastVideoArgs":
|
||||
kwargs['pipeline_config'] = PipelineConfig.from_kwargs(kwargs)
|
||||
return cls(**kwargs)
|
||||
|
||||
def check_fastvideo_args(self) -> None:
|
||||
"""Validate inference arguments for consistency"""
|
||||
if not self.inference_mode:
|
||||
assert self.dp_size is not -1, "dp_size must be set for training"
|
||||
assert self.dp_shards is not -1, "dp_shards must be set for training"
|
||||
assert self.sp_size is not -1, "sp_size must be set for training"
|
||||
assert self.hsdp_replicate_dim != -1, "hsdp_replicate_dim must be set for training"
|
||||
assert self.hsdp_shard_dim != -1, "hsdp_shard_dim must be set for training"
|
||||
assert self.sp_size != -1, "sp_size must be set for training"
|
||||
|
||||
if self.tp_size is -1:
|
||||
if self.tp_size == -1:
|
||||
self.tp_size = self.num_gpus
|
||||
if self.sp_size is -1:
|
||||
if self.sp_size == -1:
|
||||
self.sp_size = self.num_gpus
|
||||
if self.dp_shards is -1:
|
||||
self.dp_shards = self.num_gpus
|
||||
if self.hsdp_shard_dim == -1:
|
||||
self.hsdp_shard_dim = self.num_gpus
|
||||
|
||||
assert self.sp_size <= self.num_gpus and self.num_gpus % self.sp_size == 0, "num_gpus must >= and be divisible by sp_size"
|
||||
assert self.dp_size <= self.num_gpus and self.num_gpus % self.dp_size == 0, "num_gpus must >= and be divisible by dp_size"
|
||||
assert self.dp_shards <= self.num_gpus and self.num_gpus % self.dp_shards == 0, "num_gpus must >= and be divisible by dp_shards"
|
||||
assert self.hsdp_replicate_dim <= self.num_gpus and self.num_gpus % self.hsdp_replicate_dim == 0, "num_gpus must >= and be divisible by hsdp_replicate_dim"
|
||||
assert self.hsdp_shard_dim <= self.num_gpus and self.num_gpus % self.hsdp_shard_dim == 0, "num_gpus must >= and be divisible by hsdp_shard_dim"
|
||||
|
||||
if self.num_gpus < max(self.tp_size, self.sp_size):
|
||||
self.num_gpus = max(self.tp_size, self.sp_size)
|
||||
@@ -419,33 +310,17 @@ class FastVideoArgs:
|
||||
f"tp_size ({self.tp_size}) must be equal to sp_size ({self.sp_size})"
|
||||
)
|
||||
|
||||
# Validate VAE spatial parallelism with VAE tiling
|
||||
if self.vae_sp and not self.vae_tiling:
|
||||
raise ValueError(
|
||||
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
|
||||
)
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.text_encoder_precisions):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
|
||||
)
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
)
|
||||
|
||||
if len(self.preprocess_text_funcs) != len(self.postprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
)
|
||||
|
||||
if self.enable_torch_compile and self.num_gpus > 1:
|
||||
logger.warning(
|
||||
"Currently torch compile does not work with multi-gpu. Setting enable_torch_compile to False"
|
||||
)
|
||||
self.enable_torch_compile = False
|
||||
|
||||
if self.pipeline_config is None:
|
||||
raise ValueError("pipeline_config is not set in FastVideoArgs")
|
||||
|
||||
self.pipeline_config.check_pipeline_config()
|
||||
|
||||
|
||||
_current_fastvideo_args = None
|
||||
|
||||
@@ -519,21 +394,22 @@ class TrainingArgs(FastVideoArgs):
|
||||
# text encoder & vae & diffusion model
|
||||
pretrained_model_name_or_path: str = ""
|
||||
dit_model_name_or_path: str = ""
|
||||
cache_dir: str = ""
|
||||
|
||||
# diffusion setting
|
||||
ema_decay: float = 0.0
|
||||
ema_start_step: int = 0
|
||||
cfg: float = 0.0
|
||||
training_cfg_rate: float = 0.0
|
||||
precondition_outputs: bool = False
|
||||
|
||||
# validation & logs
|
||||
validation_prompt_dir: str = ""
|
||||
validation_dataset_file: str = ""
|
||||
validation_preprocessed_path: str = ""
|
||||
validation_sampling_steps: str = ""
|
||||
validation_guidance_scale: str = ""
|
||||
validation_steps: float = 0.0
|
||||
log_validation: bool = False
|
||||
tracker_project_name: str = ""
|
||||
wandb_run_name: str = ""
|
||||
seed: Optional[int] = None
|
||||
|
||||
# output
|
||||
@@ -541,7 +417,6 @@ class TrainingArgs(FastVideoArgs):
|
||||
checkpoints_total_limit: int = 0
|
||||
checkpointing_steps: int = 0
|
||||
resume_from_checkpoint: bool = False
|
||||
logging_dir: str = ""
|
||||
|
||||
# optimizer & scheduler
|
||||
num_train_epochs: int = 0
|
||||
@@ -582,36 +457,29 @@ class TrainingArgs(FastVideoArgs):
|
||||
# master_weight_type
|
||||
master_weight_type: str = ""
|
||||
|
||||
# For fast checking in LoRA pipeline
|
||||
training_mode: bool = True
|
||||
# VSA training decay parameters
|
||||
VSA_decay_rate: float = 0.01 # decay rate -> 0.02
|
||||
VSA_decay_interval_steps: int = 1 # decay interval steps -> 50
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
provided_args = clean_cli_args(args)
|
||||
# Get all fields from the dataclass
|
||||
attrs = [attr.name for attr in dataclasses.fields(cls)]
|
||||
|
||||
logger.info(provided_args)
|
||||
# Create a dictionary of attribute values, with defaults for missing attributes
|
||||
kwargs = {}
|
||||
for attr in attrs:
|
||||
# Handle renamed attributes or those with multiple CLI names
|
||||
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
|
||||
kwargs[attr] = args.tensor_parallel_size
|
||||
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
|
||||
kwargs[attr] = args.sequence_parallel_size
|
||||
elif attr == 'flow_shift' and hasattr(args, 'shift'):
|
||||
kwargs[attr] = args.shift
|
||||
elif attr == 'dp_size' and hasattr(args, 'data_parallel_size'):
|
||||
kwargs[attr] = args.data_parallel_size
|
||||
elif attr == 'dp_shards' and hasattr(args, 'data_parallel_shards'):
|
||||
kwargs[attr] = args.data_parallel_shards
|
||||
if attr == 'pipeline_config':
|
||||
pipeline_config = PipelineConfig.from_kwargs(provided_args)
|
||||
kwargs[attr] = pipeline_config
|
||||
# Use getattr with default value from the dataclass for potentially missing attributes
|
||||
else:
|
||||
default_value = getattr(cls, attr, None)
|
||||
value = getattr(args, attr, default_value)
|
||||
if value is not None:
|
||||
kwargs[attr] = value
|
||||
kwargs[attr] = value # type: ignore
|
||||
|
||||
return cls(**kwargs)
|
||||
return cls(**kwargs) # type: ignore
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
@@ -674,7 +542,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=int,
|
||||
default=0,
|
||||
help="Step to start EMA")
|
||||
parser.add_argument("--cfg",
|
||||
parser.add_argument("--training-cfg-rate",
|
||||
type=float,
|
||||
help="Classifier-free guidance scale")
|
||||
parser.add_argument(
|
||||
@@ -683,9 +551,12 @@ class TrainingArgs(FastVideoArgs):
|
||||
help="Whether to precondition the outputs of the model")
|
||||
|
||||
# Validation and logging
|
||||
parser.add_argument("--validation-prompt-dir",
|
||||
parser.add_argument("--validation-dataset-file",
|
||||
type=str,
|
||||
help="Directory containing validation prompts")
|
||||
help="Path to unprocessed validation dataset")
|
||||
parser.add_argument("--validation-preprocessed-path",
|
||||
type=str,
|
||||
help="Path to processed validation dataset")
|
||||
parser.add_argument("--validation-sampling-steps",
|
||||
type=str,
|
||||
help="Validation sampling steps")
|
||||
@@ -701,6 +572,9 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--tracker-project-name",
|
||||
type=str,
|
||||
help="Project name for tracking")
|
||||
parser.add_argument("--wandb-run-name",
|
||||
type=str,
|
||||
help="Run name for wandb")
|
||||
parser.add_argument("--seed",
|
||||
type=int,
|
||||
default=42,
|
||||
@@ -839,4 +713,16 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=str,
|
||||
help="Master weight type")
|
||||
|
||||
# VSA parameters for training with dense to sparse adaption
|
||||
parser.add_argument(
|
||||
"--VSA-decay-rate", # decay rate, how much sparsity you want to decay each step
|
||||
type=float,
|
||||
default=TrainingArgs.VSA_decay_rate,
|
||||
help="VSA decay rate")
|
||||
parser.add_argument(
|
||||
"--VSA-decay-interval-steps", # how many steps for training with current sparsity
|
||||
type=int,
|
||||
default=TrainingArgs.VSA_decay_interval_steps,
|
||||
help="VSA decay interval steps")
|
||||
|
||||
return parser
|
||||
|
||||
@@ -5,15 +5,16 @@ import time
|
||||
from collections import defaultdict
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
# if TYPE_CHECKING:
|
||||
from fastvideo.v1.attention import AttentionMetadata
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.v1.attention import AttentionMetadata
|
||||
from fastvideo.v1.pipelines import ForwardBatch
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -36,13 +37,13 @@ class ForwardContext:
|
||||
# attn_layers: Dict[str, Any]
|
||||
# TODO: extend to support per-layer dynamic forward context
|
||||
attn_metadata: "AttentionMetadata" # set dynamically for each forward pass
|
||||
forward_batch: Optional[ForwardBatch] = None
|
||||
forward_batch: Optional["ForwardBatch"] = None
|
||||
|
||||
|
||||
_forward_context: Optional[ForwardContext] = None
|
||||
_forward_context: Optional["ForwardContext"] = None
|
||||
|
||||
|
||||
def get_forward_context() -> ForwardContext:
|
||||
def get_forward_context() -> "ForwardContext":
|
||||
"""Get the current forward context."""
|
||||
assert _forward_context is not None, (
|
||||
"Forward context is not set. "
|
||||
@@ -54,7 +55,7 @@ def get_forward_context() -> ForwardContext:
|
||||
@contextmanager
|
||||
def set_forward_context(current_timestep,
|
||||
attn_metadata,
|
||||
forward_batch: Optional[ForwardBatch] = None,
|
||||
forward_batch: Optional["ForwardBatch"] = None,
|
||||
fastvideo_args: Optional[FastVideoArgs] = None):
|
||||
"""A context manager that stores the current forward context,
|
||||
can be attention metadata, etc.
|
||||
|
||||
@@ -5,6 +5,8 @@ from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
from torch.distributed.tensor import DTensor
|
||||
|
||||
from fastvideo.v1.layers.custom_op import CustomOp
|
||||
|
||||
@@ -69,7 +71,12 @@ class RMSNorm(CustomOp):
|
||||
x = x * torch.rsqrt(variance + self.variance_epsilon)
|
||||
x = x.to(orig_dtype)
|
||||
if self.has_weight:
|
||||
x = x * self.weight
|
||||
# TODO(wenxuan): When using CPU offload, FSDP has a bug that doesn't unwrap DTensor in final_layer_norm.
|
||||
# Report this
|
||||
if isinstance(self.weight, DTensor):
|
||||
x = x * self.weight.to(x.device).full_tensor()
|
||||
else:
|
||||
x = x * self.weight
|
||||
if residual is None:
|
||||
return x
|
||||
else:
|
||||
@@ -95,6 +102,22 @@ class ScaleResidual(nn.Module):
|
||||
return residual + x * gate
|
||||
|
||||
|
||||
# adapted from Diffusers: https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/normalization.py
|
||||
# NOTE(will): Needed to match behavior of diffusers and wan2.1 even while using
|
||||
# FSDP's MixedPrecisionPolicy
|
||||
class FP32LayerNorm(nn.LayerNorm):
|
||||
|
||||
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
|
||||
origin_dtype = inputs.dtype
|
||||
return F.layer_norm(
|
||||
inputs.float(),
|
||||
self.normalized_shape,
|
||||
self.weight.float() if self.weight is not None else None,
|
||||
self.bias.float() if self.bias is not None else None,
|
||||
self.eps,
|
||||
).to(origin_dtype)
|
||||
|
||||
|
||||
class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
"""
|
||||
Fused operation that combines:
|
||||
@@ -112,6 +135,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
eps: float = 1e-6,
|
||||
elementwise_affine: bool = False,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
compute_dtype: torch.dtype | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
@@ -121,10 +145,15 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
elif norm_type == "layer":
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
if compute_dtype == torch.float32:
|
||||
self.norm = FP32LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps)
|
||||
else:
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
else:
|
||||
raise NotImplementedError(f"Norm type {norm_type} not implemented")
|
||||
|
||||
@@ -163,18 +192,25 @@ class LayerNormScaleShift(nn.Module):
|
||||
eps: float = 1e-6,
|
||||
elementwise_affine: bool = False,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
compute_dtype: torch.dtype | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.compute_dtype = compute_dtype
|
||||
if norm_type == "rms":
|
||||
self.norm = RMSNorm(hidden_size,
|
||||
has_weight=elementwise_affine,
|
||||
eps=eps)
|
||||
elif norm_type == "layer":
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
if self.compute_dtype == torch.float32:
|
||||
self.norm = FP32LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps)
|
||||
else:
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
else:
|
||||
raise NotImplementedError(f"Norm type {norm_type} not implemented")
|
||||
|
||||
@@ -182,4 +218,7 @@ class LayerNormScaleShift(nn.Module):
|
||||
scale: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply ln followed by scale and shift in a single fused operation."""
|
||||
normalized = self.norm(x)
|
||||
return normalized * (1.0 + scale) + shift
|
||||
if self.compute_dtype == torch.float32:
|
||||
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
|
||||
else:
|
||||
return normalized * (1.0 + scale) + shift
|
||||
|
||||
@@ -114,7 +114,7 @@ def _info(logger: Logger,
|
||||
|
||||
if (main_process_only and is_main_process) or (local_main_process_only
|
||||
and is_local_main_process):
|
||||
logger.log(logging.INFO, msg, *args, **kwargs)
|
||||
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
|
||||
|
||||
global _warned_local_main_process, _warned_main_process
|
||||
|
||||
@@ -134,7 +134,7 @@ def _info(logger: Logger,
|
||||
_warned_main_process = True
|
||||
|
||||
if not main_process_only and not local_main_process_only:
|
||||
logger.log(logging.INFO, msg, *args, **kwargs)
|
||||
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
|
||||
|
||||
|
||||
class _FastvideoLogger(Logger):
|
||||
|
||||
@@ -6,7 +6,7 @@ import torch
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
# TODO
|
||||
@@ -14,12 +14,13 @@ class BaseDiT(nn.Module, ABC):
|
||||
_fsdp_shard_conditions: list = []
|
||||
_compile_conditions: list = []
|
||||
_param_names_mapping: dict
|
||||
_reverse_param_names_mapping: dict
|
||||
hidden_size: int
|
||||
num_attention_heads: int
|
||||
num_channels_latents: int
|
||||
# always supports torch_sdpa
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = DiTConfig()._supported_attention_backends
|
||||
AttentionBackendEnum, ...] = DiTConfig()._supported_attention_backends
|
||||
|
||||
def __init_subclass__(cls) -> None:
|
||||
required_class_attrs = [
|
||||
@@ -65,7 +66,7 @@ class BaseDiT(nn.Module, ABC):
|
||||
)
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
|
||||
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
|
||||
@@ -78,6 +79,7 @@ class CachableDiT(BaseDiT):
|
||||
# These are required class attributes that should be overridden by concrete implementations
|
||||
_fsdp_shard_conditions = []
|
||||
_param_names_mapping = {}
|
||||
_reverse_param_names_mapping = {}
|
||||
_lora_param_names_mapping: dict = {}
|
||||
# Ensure these instance attributes are properly defined in subclasses
|
||||
hidden_size: int
|
||||
@@ -85,7 +87,7 @@ class CachableDiT(BaseDiT):
|
||||
num_channels_latents: int
|
||||
# always supports torch_sdpa
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = DiTConfig()._supported_attention_backends
|
||||
AttentionBackendEnum, ...] = DiTConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, config: DiTConfig, **kwargs) -> None:
|
||||
super().__init__(config, **kwargs)
|
||||
|
||||
@@ -23,7 +23,7 @@ from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
|
||||
unpatchify)
|
||||
from fastvideo.v1.models.dits.base import CachableDiT
|
||||
from fastvideo.v1.models.utils import modulate
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class HunyuanRMSNorm(nn.Module):
|
||||
@@ -96,7 +96,8 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
@@ -303,7 +304,8 @@ class MMSingleStreamBlock(nn.Module):
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
@@ -440,6 +442,8 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
|
||||
_supported_attention_backends = HunyuanVideoConfig(
|
||||
)._supported_attention_backends
|
||||
_param_names_mapping = HunyuanVideoConfig()._param_names_mapping
|
||||
_reverse_param_names_mapping = HunyuanVideoConfig(
|
||||
)._reverse_param_names_mapping
|
||||
_lora_param_names_mapping = HunyuanVideoConfig()._lora_param_names_mapping
|
||||
|
||||
def __init__(self, config: HunyuanVideoConfig, hf_config: dict[str, Any]):
|
||||
@@ -876,8 +880,8 @@ class IndividualTokenRefinerBlock(nn.Module):
|
||||
num_heads=num_attention_heads,
|
||||
head_size=hidden_size // num_attention_heads,
|
||||
# TODO: remove hardcode; remove STA
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA),
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA),
|
||||
)
|
||||
|
||||
def forward(self, x, c):
|
||||
|
||||
@@ -26,7 +26,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
|
||||
get_rotary_pos_embed)
|
||||
from fastvideo.v1.layers.visual_embedding import TimestepEmbedder
|
||||
from fastvideo.v1.models.dits.base import BaseDiT
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class PatchEmbed2D(nn.Module):
|
||||
@@ -139,16 +139,17 @@ class StepVideoRMSNorm(nn.Module):
|
||||
|
||||
class SelfAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
hidden_dim,
|
||||
head_dim,
|
||||
rope_split: Tuple[int, int, int] = (64, 32, 32),
|
||||
bias: bool = False,
|
||||
with_rope: bool = True,
|
||||
with_qk_norm: bool = True,
|
||||
attn_type: str = "torch",
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA)):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_dim,
|
||||
head_dim,
|
||||
rope_split: Tuple[int, int, int] = (64, 32, 32),
|
||||
bias: bool = False,
|
||||
with_rope: bool = True,
|
||||
with_qk_norm: bool = True,
|
||||
attn_type: str = "torch",
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA)):
|
||||
super().__init__()
|
||||
self.head_dim = head_dim
|
||||
self.hidden_dim = hidden_dim
|
||||
@@ -257,7 +258,8 @@ class CrossAttention(nn.Module):
|
||||
head_dim,
|
||||
bias=False,
|
||||
with_qk_norm=True,
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA)
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.head_dim = head_dim
|
||||
@@ -453,11 +455,10 @@ class StepVideoTransformerBlock(nn.Module):
|
||||
|
||||
class StepVideoModel(BaseDiT):
|
||||
# (Optional) Keep the same attribute for compatibility with splitting, etc.
|
||||
_fsdp_shard_conditions = [
|
||||
lambda n, m: "transformer_blocks" in n and n.split(".")[-1].isdigit(),
|
||||
# lambda n, m: "pos_embed" in n # If needed for the patch embedding.
|
||||
]
|
||||
_fsdp_shard_conditions = StepVideoConfig()._fsdp_shard_conditions
|
||||
_param_names_mapping = StepVideoConfig()._param_names_mapping
|
||||
_reverse_param_names_mapping = StepVideoConfig(
|
||||
)._reverse_param_names_mapping
|
||||
_lora_param_names_mapping = StepVideoConfig()._lora_param_names_mapping
|
||||
_supported_attention_backends = StepVideoConfig(
|
||||
)._supported_attention_backends
|
||||
|
||||
@@ -14,8 +14,8 @@ from fastvideo.v1.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.v1.configs.sample.wan import WanTeaCacheParams
|
||||
from fastvideo.v1.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.v1.forward_context import get_forward_context
|
||||
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, RMSNorm,
|
||||
ScaleResidual,
|
||||
from fastvideo.v1.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
|
||||
RMSNorm, ScaleResidual,
|
||||
ScaleResidualLayerNormScaleShift)
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
# from torch.nn import RMSNorm
|
||||
@@ -26,7 +26,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
|
||||
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
|
||||
PatchEmbed, TimestepEmbedder)
|
||||
from fastvideo.v1.models.dits.base import CachableDiT
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class WanImageEmbedding(torch.nn.Module):
|
||||
@@ -34,9 +34,9 @@ class WanImageEmbedding(torch.nn.Module):
|
||||
def __init__(self, in_features: int, out_features: int):
|
||||
super().__init__()
|
||||
|
||||
self.norm1 = nn.LayerNorm(in_features)
|
||||
self.norm1 = FP32LayerNorm(in_features)
|
||||
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
|
||||
self.norm2 = nn.LayerNorm(out_features)
|
||||
self.norm2 = FP32LayerNorm(out_features)
|
||||
|
||||
def forward(self,
|
||||
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
@@ -125,8 +125,8 @@ class WanSelfAttention(nn.Module):
|
||||
dropout_rate=0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA))
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA))
|
||||
|
||||
def forward(self, x: torch.Tensor, context: torch.Tensor,
|
||||
context_lens: int):
|
||||
@@ -174,7 +174,8 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
eps=1e-6,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None
|
||||
) -> None:
|
||||
super().__init__(dim, num_heads, window_size, qk_norm, eps,
|
||||
supported_attention_backends)
|
||||
@@ -216,21 +217,22 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
|
||||
class WanTransformerBlock(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
prefix: str = ""):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -261,7 +263,8 @@ class WanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
@@ -281,7 +284,8 @@ class WanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -358,21 +362,22 @@ class WanTransformerBlock(nn.Module):
|
||||
|
||||
class WanTransformerBlock_VSA(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
prefix: str = ""):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -404,7 +409,8 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
@@ -424,7 +430,8 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -511,6 +518,7 @@ class WanTransformer3DModel(CachableDiT):
|
||||
_supported_attention_backends = WanVideoConfig(
|
||||
)._supported_attention_backends
|
||||
_param_names_mapping = WanVideoConfig()._param_names_mapping
|
||||
_reverse_param_names_mapping = WanVideoConfig()._reverse_param_names_mapping
|
||||
_lora_param_names_mapping = WanVideoConfig()._lora_param_names_mapping
|
||||
|
||||
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
|
||||
@@ -561,7 +569,8 @@ class WanTransformer3DModel(CachableDiT):
|
||||
norm_type="layer",
|
||||
eps=config.eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
@@ -569,6 +578,17 @@ class WanTransformer3DModel(CachableDiT):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
# For type checking
|
||||
self.previous_e0_even = None
|
||||
self.previous_e0_odd = None
|
||||
self.previous_residual_even = None
|
||||
self.previous_residual_odd = None
|
||||
self.is_even = True
|
||||
self.should_calc_even = True
|
||||
self.should_calc_odd = True
|
||||
self.accumulated_rel_l1_distance_even = 0
|
||||
self.accumulated_rel_l1_distance_odd = 0
|
||||
self.cnt = 0
|
||||
self.__post_init__()
|
||||
|
||||
def forward(self,
|
||||
@@ -655,7 +675,7 @@ class WanTransformer3DModel(CachableDiT):
|
||||
# 5. Output norm, projection & unpatchify
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
|
||||
dim=1)
|
||||
hidden_states = self.norm_out(hidden_states.float(), shift, scale)
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Optional, Tuple
|
||||
from dataclasses import field
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
@@ -8,16 +9,22 @@ from torch import nn
|
||||
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
|
||||
ImageEncoderConfig,
|
||||
TextEncoderConfig)
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class TextEncoder(nn.Module, ABC):
|
||||
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
|
||||
_stacked_params_mapping: List[Tuple[str, str,
|
||||
str]] = field(default_factory=list)
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = TextEncoderConfig()._supported_attention_backends
|
||||
AttentionBackendEnum,
|
||||
...] = TextEncoderConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, config: TextEncoderConfig) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self._fsdp_shard_conditions = config._fsdp_shard_conditions
|
||||
self._stacked_params_mapping = config.arch_config.stacked_params_mapping
|
||||
if not self.supported_attention_backends:
|
||||
raise ValueError(
|
||||
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
|
||||
@@ -34,13 +41,14 @@ class TextEncoder(nn.Module, ABC):
|
||||
pass
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
|
||||
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
|
||||
class ImageEncoder(nn.Module, ABC):
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = ImageEncoderConfig()._supported_attention_backends
|
||||
AttentionBackendEnum,
|
||||
...] = ImageEncoderConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, config: ImageEncoderConfig) -> None:
|
||||
super().__init__()
|
||||
@@ -56,5 +64,5 @@ class ImageEncoder(nn.Module, ABC):
|
||||
pass
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
|
||||
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
@@ -596,12 +596,7 @@ class CLIPVisionModel(ImageEncoder):
|
||||
# ref: https://github.com/vllm-project/vllm/pull/7186#discussion_r1734163986
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
("qkv_proj", "q_proj", "q"),
|
||||
("qkv_proj", "k_proj", "k"),
|
||||
("qkv_proj", "v_proj", "v"),
|
||||
]
|
||||
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: Set[str] = set()
|
||||
layer_count = len(self.vision_model.encoder.layers)
|
||||
@@ -620,7 +615,8 @@ class CLIPVisionModel(ImageEncoder):
|
||||
if layer_idx >= layer_count:
|
||||
continue
|
||||
|
||||
for (param_name, weight_name, shard_id) in stacked_params_mapping:
|
||||
for (param_name, weight_name,
|
||||
shard_id) in self.config.arch_config.stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
|
||||
@@ -369,14 +369,7 @@ class LlamaModel(TextEncoder):
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str,
|
||||
torch.Tensor]]) -> Set[str]:
|
||||
stacked_params_mapping = [
|
||||
# (param_name, shard_name, shard_id)
|
||||
(".qkv_proj", ".q_proj", "q"),
|
||||
(".qkv_proj", ".k_proj", "k"),
|
||||
(".qkv_proj", ".v_proj", "v"),
|
||||
(".gate_up_proj", ".gate_proj", 0),
|
||||
(".gate_up_proj", ".up_proj", 1),
|
||||
]
|
||||
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: Set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
@@ -406,7 +399,7 @@ class LlamaModel(TextEncoder):
|
||||
continue
|
||||
else:
|
||||
name = kv_scale_name
|
||||
for param_name, weight_name, shard_id in stacked_params_mapping:
|
||||
for param_name, weight_name, shard_id in self.config.arch_config.stacked_params_mapping:
|
||||
if weight_name not in name:
|
||||
continue
|
||||
name = name.replace(weight_name, param_name)
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user