Compare commits
13
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ea8c154176 | ||
|
|
580d6dfe1f | ||
|
|
344e43006a | ||
|
|
c5155b256e | ||
|
|
e005c7f3ac | ||
|
|
ff5a79ef60 | ||
|
|
ab01dc4ba5 | ||
|
|
285a950c1b | ||
|
|
46a0a85d85 | ||
|
|
4aeabbc629 | ||
|
|
949bb5c835 | ||
|
|
aab74c1271 | ||
|
|
f89d86944f |
+95
-13
@@ -2,25 +2,26 @@ env:
|
||||
IMAGE_VERSION: "py3.12-latest"
|
||||
|
||||
steps:
|
||||
- block: "Start Build"
|
||||
blocked_state: "running"
|
||||
prompt: "Approve build?"
|
||||
- label: "pre-commit"
|
||||
command: ".buildkite/scripts/pre_commit.sh"
|
||||
agents:
|
||||
queue: "default"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
|
||||
- wait
|
||||
|
||||
- label: "Trigger Tests"
|
||||
command: |
|
||||
echo "Current working directory: $(pwd)"
|
||||
echo "Current branch:"
|
||||
git branch --show-current
|
||||
echo "Full diff:"
|
||||
git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD
|
||||
plugins:
|
||||
- monorepo-diff#v1.4.0:
|
||||
diff: "git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD"
|
||||
watch:
|
||||
- path:
|
||||
- "fastvideo/v1/models/encoders/**"
|
||||
- "fastvideo/v1/models/loaders/**"
|
||||
- "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"
|
||||
@@ -31,8 +32,10 @@ steps:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/models/vaes/**"
|
||||
- "fastvideo/v1/models/loaders/**"
|
||||
- "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"
|
||||
@@ -43,10 +46,12 @@ steps:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/models/dits/**"
|
||||
- "fastvideo/v1/models/loaders/**"
|
||||
- "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"
|
||||
@@ -55,7 +60,8 @@ steps:
|
||||
- TEST_TYPE=transformer
|
||||
agents:
|
||||
queue: "default"
|
||||
- path: "fastvideo/v1/**/*.py"
|
||||
- path:
|
||||
- "fastvideo/v1/**/*.py"
|
||||
config:
|
||||
command: "timeout 60m .buildkite/scripts/pr_test.sh"
|
||||
label: "SSIM Tests"
|
||||
@@ -64,3 +70,79 @@ steps:
|
||||
- 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"
|
||||
|
||||
@@ -31,6 +31,10 @@ log "Setting up Modal authentication from Buildkite secrets..."
|
||||
MODAL_TOKEN_ID=$(buildkite-agent secret get modal_token_id)
|
||||
MODAL_TOKEN_SECRET=$(buildkite-agent secret get modal_token_secret)
|
||||
|
||||
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
|
||||
|
||||
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
|
||||
|
||||
if [ -n "$MODAL_TOKEN_ID" ] && [ -n "$MODAL_TOKEN_SECRET" ]; then
|
||||
log "Retrieved Modal credentials from Buildkite secrets"
|
||||
python3 -m modal token set --token-id "$MODAL_TOKEN_ID" --token-secret "$MODAL_TOKEN_SECRET" --profile buildkite-ci --activate --verify
|
||||
@@ -54,22 +58,44 @@ if [ -z "${TEST_TYPE:-}" ]; then
|
||||
fi
|
||||
log "Test type: $TEST_TYPE"
|
||||
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT IMAGE_VERSION=$IMAGE_VERSION"
|
||||
|
||||
case "$TEST_TYPE" in
|
||||
"encoder")
|
||||
log "Running encoder tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
|
||||
;;
|
||||
"vae")
|
||||
log "Running VAE tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
|
||||
;;
|
||||
"transformer")
|
||||
log "Running transformer tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
|
||||
;;
|
||||
"ssim")
|
||||
log "Running SSIM tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
;;
|
||||
"training")
|
||||
log "Running training tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests"
|
||||
;;
|
||||
"training_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"
|
||||
|
||||
@@ -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
|
||||
+104
-34
@@ -14,13 +14,9 @@ on:
|
||||
- ".github/workflows/pr-test.yml"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
- "csrc/**"
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
custom_image:
|
||||
description: "Custom image from this repository (default: fastvideo-dev:py3.12-latest)"
|
||||
required: false
|
||||
default: "fastvideo-dev:py3.12-latest"
|
||||
type: string
|
||||
run_encoder_test:
|
||||
description: "Run encoder-test"
|
||||
required: false
|
||||
@@ -56,6 +52,16 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_precision_test_STA:
|
||||
description: "Run precision-test-STA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_precision_test_VSA:
|
||||
description: "Run precision-test-VSA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_nightly_test:
|
||||
description: "Run nightly-test"
|
||||
required: false
|
||||
@@ -65,6 +71,7 @@ on:
|
||||
env:
|
||||
PYTHONUNBUFFERED: "1"
|
||||
|
||||
|
||||
concurrency:
|
||||
group: pr-test-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
@@ -84,44 +91,69 @@ jobs:
|
||||
training-test: ${{ steps.filter.outputs.training-test }}
|
||||
training-test-VSA: ${{ steps.filter.outputs.training-test-VSA }}
|
||||
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
|
||||
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
|
||||
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
id: filter
|
||||
with:
|
||||
filters: |
|
||||
# Define reusable path patterns
|
||||
common-paths: &common-paths
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
sta-kernel-paths: &sta-kernel-paths
|
||||
- 'csrc/attn/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/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
- *common-paths
|
||||
vae-test:
|
||||
- 'fastvideo/v1/models/vaes/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/models/loader/**'
|
||||
- 'fastvideo/v1/tests/vaes/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
- *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/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
- *common-paths
|
||||
training-test:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
- *common-paths
|
||||
training-test-VSA:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
inference-test-STA:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
- *common-paths
|
||||
- *sta-kernel-paths
|
||||
precision-test-STA:
|
||||
- *common-paths
|
||||
- *sta-kernel-paths
|
||||
precision-test-VSA:
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
@@ -134,7 +166,7 @@ jobs:
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
|
||||
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:
|
||||
@@ -152,7 +184,7 @@ jobs:
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
|
||||
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:
|
||||
@@ -170,7 +202,7 @@ jobs:
|
||||
gpu_type: "NVIDIA L40S"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
|
||||
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:
|
||||
@@ -180,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:
|
||||
@@ -207,7 +238,7 @@ jobs:
|
||||
training-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
@@ -216,7 +247,7 @@ jobs:
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
|
||||
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:
|
||||
@@ -227,7 +258,7 @@ jobs:
|
||||
training-test-VSA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test-VSA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test_VSA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
@@ -236,7 +267,7 @@ jobs:
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
|
||||
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:
|
||||
@@ -247,7 +278,7 @@ jobs:
|
||||
inference-test-STA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.inference-test-STA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_inference_test_STA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
@@ -256,13 +287,51 @@ jobs:
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
|
||||
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')
|
||||
@@ -273,7 +342,7 @@ jobs:
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
|
||||
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:
|
||||
@@ -282,7 +351,8 @@ jobs:
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
runpod-cleanup:
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
|
||||
# Add other jobs to this list as you create them
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test, training-test, training-test-VSA, inference-test-STA, precision-test-STA, precision-test-VSA]
|
||||
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
@@ -299,7 +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
|
||||
|
||||
+8
-2
@@ -4,7 +4,7 @@
|
||||
|
||||
|
||||
## 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
|
||||
@@ -53,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();
|
||||
}
|
||||
|
||||
|
||||
@@ -7,6 +7,7 @@ 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)
|
||||
@@ -81,5 +81,7 @@ std = 10
|
||||
|
||||
# Run correctness check directly
|
||||
results = check_correctness(b, h, n, d, causal, mean, std, error_mode='output')
|
||||
assert results['TK vs FLEX']['avg_diff'] < 3e-6, f"Average difference: {results['TK vs FLEX']['avg_diff']} is too large"
|
||||
assert results['TK vs FLEX']['max_diff'] < 4e-2, f"Maximum difference: {results['TK vs FLEX']['max_diff']} is too large"
|
||||
print(f"Average difference: {results['TK vs FLEX']['avg_diff']}")
|
||||
print(f"Maximum difference: {results['TK vs FLEX']['max_diff']}")
|
||||
@@ -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();
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
@@ -1,7 +1,10 @@
|
||||
#!/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
|
||||
@@ -12,7 +15,7 @@ NUM_GPUS=8
|
||||
training_args=(
|
||||
--tracker_project_name "wan_i2v_finetune"
|
||||
--output_dir "$DATA_DIR/outputs/wan_i2v_finetune"
|
||||
--max_train_steps 5000
|
||||
--max_train_steps 2000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
@@ -33,8 +36,8 @@ parallel_args=(
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
--pretrained_model_name_or_path "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
@@ -56,7 +59,7 @@ validation_args=(
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 6000
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
@@ -66,7 +69,7 @@ miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--cfg 0.0
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=FV_2N_14B
|
||||
#SBATCH --job-name=i2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --qos=hao
|
||||
#SBATCH --nodes=4
|
||||
@@ -7,10 +7,10 @@
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --nodelist=fs-mbz-gpu-[400-550]
|
||||
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=4n_i2v/4n_i2v_%j.out
|
||||
#SBATCH --error=4n_i2v/4n_i2v_%j.err
|
||||
#SBATCH --output=i2v_output/i2v_%j.out
|
||||
#SBATCH --error=i2v_output/i2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
@@ -30,7 +30,9 @@ nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
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"
|
||||
@@ -38,60 +40,91 @@ echo "NODE_RANK: $NODE_RANK"
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
DATA_DIR=data/crush-smol_processed_i2v/combined_parquet_dataset
|
||||
VALIDATION_DIR=data/crush-smol_processed_i2v/validation_parquet_dataset
|
||||
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\
|
||||
--model_path Wan-AI/Wan2.1-I2V-14B-480P-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-I2V-14B-480P-Diffusers \
|
||||
--cache_dir "/home/ray/.cache"\
|
||||
--data_path "$DATA_DIR"\
|
||||
--validation_preprocessed_path "$VALIDATION_DIR"\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 16 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
--sp_size $NUM_GPUS \
|
||||
--tp_size $NUM_GPUS \
|
||||
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES \
|
||||
--hsdp_shard_dim $NUM_GPUS \
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 10\
|
||||
--gradient_accumulation_steps=2\
|
||||
--max_train_steps=10000 \
|
||||
--learning_rate=5e-5\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=11000 \
|
||||
--validation_steps 100\
|
||||
--validation_sampling_steps "40" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune_2n"\
|
||||
--tracker_project_name wan_i2v_finetune \
|
||||
--num_height 480 \
|
||||
--num_width 832 \
|
||||
--num_frames 77 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--weight_decay 1e-4 \
|
||||
--not_apply_cfg_solver \
|
||||
--dit_precision "fp32" \
|
||||
--max_grad_norm 1.0
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
+2
-1
@@ -1,4 +1,5 @@
|
||||
# export WANDB_MODE="offline"
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
@@ -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"
|
||||
@@ -1,3 +1,5 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
@@ -16,7 +18,7 @@ training_args=(
|
||||
--max_train_steps 5000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--gradient_accumulation_steps 8
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
@@ -25,11 +27,11 @@ training_args=(
|
||||
|
||||
# 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 \
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size $NUM_GPUS
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
@@ -48,17 +50,17 @@ dataset_args=(
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_preprocessed_path $VALIDATION_DIR
|
||||
--validation_steps 100
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--learning_rate 5e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 6000
|
||||
--weight_decay 0.01
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
@@ -67,7 +69,7 @@ miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--cfg 0.0
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=FV_2N_14B
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --qos=hao
|
||||
#SBATCH --nodes=1
|
||||
@@ -7,10 +7,10 @@
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --nodelist=fs-mbz-gpu-[400-550]
|
||||
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=4n_i2v/4n_i2v_%j.out
|
||||
#SBATCH --error=4n_i2v/4n_i2v_%j.err
|
||||
#SBATCH --output=t2v_output/t2v_%j.out
|
||||
#SBATCH --error=t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
@@ -30,69 +30,98 @@ nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
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
|
||||
|
||||
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=8
|
||||
# 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_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\
|
||||
--model_path $MODEL_PATH \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path $MODEL_PATH \
|
||||
--cache_dir "/home/ray/.cache"\
|
||||
--data_path "$DATA_DIR"\
|
||||
--validation_preprocessed_path "$VALIDATION_DIR"\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 8 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
--sp_size $NUM_GPUS \
|
||||
--tp_size $NUM_GPUS \
|
||||
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES \
|
||||
--hsdp_shard_dim $NUM_GPUS \
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 10\
|
||||
--gradient_accumulation_steps=1\
|
||||
--max_train_steps=10000 \
|
||||
--learning_rate=5e-5\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=11000 \
|
||||
--validation_steps 100\
|
||||
--validation_sampling_steps "40" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune_2n"\
|
||||
--tracker_project_name wan_i2v_finetune \
|
||||
--num_height 480 \
|
||||
--num_width 832 \
|
||||
--num_frames 77 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--weight_decay 1e-4 \
|
||||
--not_apply_cfg_solver \
|
||||
--dit_precision "fp32" \
|
||||
--max_grad_norm 1.0
|
||||
fastvideo/v1/training/wan_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -1,13 +0,0 @@
|
||||
{
|
||||
"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": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-034.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
+2
-1
@@ -1,4 +1,5 @@
|
||||
# export WANDB_MODE="offline"
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
@@ -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
|
||||
|
||||
@@ -12,6 +12,7 @@ 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[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -11,7 +11,7 @@ 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_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
|
||||
|
||||
@@ -148,8 +148,8 @@ def main() -> None:
|
||||
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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -67,14 +68,18 @@ def main() -> None:
|
||||
|
||||
# Create DataLoader with proper settings
|
||||
dataset, dataloader = build_parquet_map_style_dataloader(
|
||||
args.path, args.batch_size, args.num_data_workers)
|
||||
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:
|
||||
# First pass - record latent sums
|
||||
first_pass_sums = []
|
||||
for i, (latents, embeddings, masks,
|
||||
caption_text) in enumerate(dataloader):
|
||||
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)
|
||||
@@ -100,14 +105,18 @@ def main() -> None:
|
||||
|
||||
# Recreate dataloader and load state
|
||||
dataset, dataloader = build_parquet_map_style_dataloader(
|
||||
args.path, args.batch_size, args.num_data_workers)
|
||||
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,
|
||||
caption_text) in enumerate(dataloader):
|
||||
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",
|
||||
@@ -116,11 +125,16 @@ def main() -> None:
|
||||
break
|
||||
|
||||
dataset, dataloader = build_parquet_map_style_dataloader(
|
||||
args.path, args.batch_size, args.num_data_workers)
|
||||
args.path,
|
||||
args.batch_size,
|
||||
parquet_schema=pyarrow_schema_t2v,
|
||||
num_data_workers=args.num_data_workers)
|
||||
|
||||
# Second pass - verify latent sums match
|
||||
second_pass_sums = []
|
||||
for i, (latents, embeddings, masks) in enumerate(dataloader):
|
||||
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)",
|
||||
@@ -144,14 +158,15 @@ def main() -> None:
|
||||
total_samples = 0
|
||||
total_batches = 0
|
||||
for _ in range(args.num_epoch):
|
||||
for i, (latents, embeddings, masks,
|
||||
caption_text) 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)
|
||||
|
||||
@@ -185,9 +185,6 @@ 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", "latent"), "text_embedding", "clip_feature",
|
||||
"first_frame_latent", "pil_image"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
@@ -204,10 +201,6 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
self.path = path
|
||||
self.cfg_rate = cfg_rate
|
||||
self.parquet_schema = parquet_schema
|
||||
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"
|
||||
)
|
||||
logger.info("Initializing LatentsParquetMapStyleDataset with path: %s",
|
||||
path)
|
||||
self.parquet_files, self.lengths = get_parquet_files_and_length(path)
|
||||
@@ -243,7 +236,8 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
|
||||
batch = collate_rows_from_parquet_schema([row_dict],
|
||||
self.parquet_schema,
|
||||
self.text_padding_length)
|
||||
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']
|
||||
@@ -265,11 +259,10 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
for idx in indices
|
||||
]
|
||||
|
||||
# all_latents, all_embs, all_masks, caption_text, all_extra_latents, all_infos = collate_latents_embs_masks(
|
||||
# rows, self.text_padding_length, self.keys)
|
||||
# return all_latents, all_embs, all_masks, caption_text, all_extra_latents, all_infos
|
||||
batch = collate_rows_from_parquet_schema(rows, self.parquet_schema,
|
||||
self.text_padding_length)
|
||||
batch = collate_rows_from_parquet_schema(rows,
|
||||
self.parquet_schema,
|
||||
self.text_padding_length,
|
||||
cfg_rate=self.cfg_rate)
|
||||
return batch
|
||||
|
||||
def __len__(self):
|
||||
|
||||
@@ -411,6 +411,29 @@ 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
|
||||
@@ -461,30 +484,30 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
self.text_encoding_stage = TextEncodingStage(
|
||||
tokenizer=tokenizer,
|
||||
text_max_length=args.text_max_length,
|
||||
cfg_rate=args.cfg)
|
||||
cfg_rate=args.training_cfg_rate)
|
||||
|
||||
def _load_raw_data(self) -> List[Dict]:
|
||||
"""Load raw data from JSON files."""
|
||||
all_data = []
|
||||
|
||||
# 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]
|
||||
|
||||
# Process each folder-annotation pair
|
||||
for folder, annotation_file in folder_anno_pairs:
|
||||
with open(annotation_file) as f:
|
||||
data_items = json.load(f)
|
||||
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"])
|
||||
# Update paths with folder prefix
|
||||
for item in data_items:
|
||||
item["path"] = opj(folder, item["path"])
|
||||
|
||||
all_data.extend(data_items)
|
||||
|
||||
return all_data[self.start_idx:]
|
||||
return data_items
|
||||
|
||||
def _process_metadata(self) -> List[PreprocessBatch]:
|
||||
"""Process the raw metadata through all filtering stages."""
|
||||
|
||||
@@ -1,12 +1,9 @@
|
||||
from typing import Any, Dict, List
|
||||
import random
|
||||
from typing import Any, Dict, List, cast
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
|
||||
"""
|
||||
@@ -24,7 +21,7 @@ def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
|
||||
return t[:padding_length], torch.ones(padding_length)
|
||||
|
||||
|
||||
def get_torch_tensors_from_row_dict(row_dict, keys) -> Dict[str, Any]:
|
||||
def get_torch_tensors_from_row_dict(row_dict, keys, cfg_rate) -> Dict[str, Any]:
|
||||
"""
|
||||
Get the latents and prompts from a row dictionary.
|
||||
"""
|
||||
@@ -42,70 +39,45 @@ def get_torch_tensors_from_row_dict(row_dict, keys) -> Dict[str, Any]:
|
||||
if shape is None or bytes is None:
|
||||
raise ValueError(f"Key {key} not found in row_dict")
|
||||
else:
|
||||
try:
|
||||
shape = row_dict[f"{key}_shape"]
|
||||
bytes = row_dict[f"{key}_bytes"]
|
||||
except KeyError:
|
||||
continue
|
||||
shape = row_dict[f"{key}_shape"]
|
||||
bytes = row_dict[f"{key}_bytes"]
|
||||
|
||||
# TODO (peiyuan): read precision
|
||||
if len(bytes) == 0:
|
||||
return_dict[key] = torch.zeros(0, dtype=torch.bfloat16)
|
||||
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
|
||||
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
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, List[str], Dict[str, Any],
|
||||
List[Dict[str, Any]]]:
|
||||
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 = []
|
||||
all_clip_features = []
|
||||
all_first_frame_latents = []
|
||||
all_pil_images = []
|
||||
all_infos = []
|
||||
caption_text = []
|
||||
# Process each row individually
|
||||
for i, row in enumerate(batch_to_process):
|
||||
# Get info from row
|
||||
info_keys = [
|
||||
"caption", "file_name", "media_type", "width", "height",
|
||||
"num_frames", "duration_sec", "fps"
|
||||
]
|
||||
info = {}
|
||||
for key in info_keys:
|
||||
if key in row:
|
||||
info[key] = row[key]
|
||||
else:
|
||||
info[key] = ""
|
||||
info["prompt"] = info["caption"]
|
||||
|
||||
# Get tensors from row
|
||||
data = get_torch_tensors_from_row_dict(row, keys)
|
||||
data = get_torch_tensors_from_row_dict(row, keys, cfg_rate)
|
||||
latents, emb = data["vae_latent"], data["text_embedding"]
|
||||
clip_feature = data.get("clip_feature", None)
|
||||
first_frame_latent = data.get("first_frame_latent", None)
|
||||
pil_image = data.get("pil_image", None)
|
||||
|
||||
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)
|
||||
all_clip_features.append(clip_feature)
|
||||
all_first_frame_latents.append(first_frame_latent)
|
||||
all_pil_images.append(pil_image)
|
||||
all_infos.append(info)
|
||||
# TODO(py): remove this once we fix preprocess
|
||||
try:
|
||||
caption_text.append(row["prompt"])
|
||||
@@ -116,17 +88,14 @@ def collate_latents_embs_masks(
|
||||
all_latents = torch.stack(all_latents)
|
||||
all_embs = torch.stack(all_embs)
|
||||
all_masks = torch.stack(all_masks)
|
||||
all_extra_latents = {
|
||||
"clip_feature": torch.stack(all_clip_features),
|
||||
"first_frame_latent": torch.stack(all_first_frame_latents),
|
||||
"pil_image": all_pil_images,
|
||||
}
|
||||
|
||||
return all_latents, all_embs, all_masks, caption_text, all_extra_latents, all_infos
|
||||
return all_latents, all_embs, all_masks, caption_text
|
||||
|
||||
|
||||
def collate_rows_from_parquet_schema(rows, parquet_schema,
|
||||
text_padding_length) -> Dict[str, Any]:
|
||||
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.
|
||||
@@ -139,10 +108,10 @@ def collate_rows_from_parquet_schema(rows, parquet_schema,
|
||||
Dict containing batched tensors and metadata
|
||||
"""
|
||||
if not rows:
|
||||
return {}
|
||||
return cast(Dict[str, Any], {})
|
||||
|
||||
# Initialize containers for different data types
|
||||
batch_data = {}
|
||||
batch_data: Dict[str, Any] = {}
|
||||
|
||||
# Get tensor and metadata field names from schema (fields ending with '_bytes')
|
||||
tensor_fields = []
|
||||
@@ -159,7 +128,7 @@ def collate_rows_from_parquet_schema(rows, parquet_schema,
|
||||
# Only add actual metadata fields, not the shape/dtype helper fields
|
||||
metadata_fields.append(field)
|
||||
|
||||
# Process each tensor field efficiently
|
||||
# Process each tensor field
|
||||
for tensor_name in tensor_fields:
|
||||
tensor_list = []
|
||||
|
||||
@@ -169,9 +138,6 @@ def collate_rows_from_parquet_schema(rows, parquet_schema,
|
||||
bytes_key = f"{tensor_name}_bytes"
|
||||
|
||||
if shape_key in row and bytes_key in row:
|
||||
# logger.info("row: %s", row)
|
||||
# logger.info("shape_key: %s", shape_key)
|
||||
# logger.info("bytes_key: %s", bytes_key)
|
||||
shape = row[shape_key]
|
||||
bytes_data = row[bytes_key]
|
||||
|
||||
@@ -179,11 +145,12 @@ def collate_rows_from_parquet_schema(rows, parquet_schema,
|
||||
tensor = torch.zeros(0, dtype=torch.bfloat16)
|
||||
else:
|
||||
# Convert bytes to tensor using float32 as default
|
||||
# logger.info("len(bytes_data): %s", len(bytes_data))
|
||||
# logger.info("shape: %s", shape)
|
||||
|
||||
data = np.frombuffer(
|
||||
bytes_data, dtype=np.float32).reshape(shape).copy()
|
||||
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
|
||||
@@ -196,38 +163,44 @@ def collate_rows_from_parquet_schema(rows, parquet_schema,
|
||||
tensor_list.append(torch.zeros(0, dtype=torch.bfloat16))
|
||||
|
||||
# Stack tensors with special handling for text embeddings
|
||||
if tensor_list:
|
||||
if tensor_name == 'text_embedding':
|
||||
# Handle text embeddings with padding
|
||||
padded_tensors = []
|
||||
attention_masks = []
|
||||
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))
|
||||
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 other tensors directly, handling None values
|
||||
valid_tensors = [
|
||||
t for t in tensor_list if t is not None and t.numel() > 0
|
||||
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
|
||||
]
|
||||
if valid_tensors:
|
||||
batch_data[tensor_name] = torch.stack(valid_tensors)
|
||||
elif tensor_list: # All tensors are empty but exist
|
||||
batch_data[tensor_name] = torch.stack(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 efficiently into info_list
|
||||
# Process metadata fields into info_list
|
||||
info_list = []
|
||||
for row in rows:
|
||||
info = {}
|
||||
|
||||
@@ -19,7 +19,6 @@ class ValidationDataset(torch.utils.data.IterableDataset):
|
||||
|
||||
self.filename = pathlib.Path(filename)
|
||||
# get directory of filename
|
||||
# TODO(will)
|
||||
self.dir = os.path.abspath(self.filename.parent)
|
||||
|
||||
if not self.filename.exists():
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -58,8 +58,10 @@ class FastVideoArgs:
|
||||
|
||||
output_type: str = "pil"
|
||||
|
||||
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
|
||||
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
@@ -208,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",
|
||||
@@ -216,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,
|
||||
@@ -384,7 +398,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
# 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
|
||||
@@ -528,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(
|
||||
|
||||
@@ -6,6 +6,7 @@ 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
|
||||
|
||||
@@ -70,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:
|
||||
|
||||
@@ -14,6 +14,7 @@ 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
|
||||
@@ -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
|
||||
|
||||
@@ -442,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]):
|
||||
|
||||
@@ -455,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
|
||||
|
||||
@@ -25,12 +25,9 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
|
||||
get_rotary_pos_embed)
|
||||
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
|
||||
PatchEmbed, TimestepEmbedder)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.dits.base import CachableDiT
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanImageEmbedding(torch.nn.Module):
|
||||
|
||||
@@ -521,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,
|
||||
|
||||
@@ -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
|
||||
@@ -12,6 +13,9 @@ 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[
|
||||
AttentionBackendEnum,
|
||||
...] = TextEncoderConfig()._supported_attention_backends
|
||||
@@ -19,6 +23,8 @@ class TextEncoder(nn.Module, ABC):
|
||||
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"
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -494,7 +494,7 @@ class T5Stack(nn.Module):
|
||||
attention_mask=attention_mask,
|
||||
attn_metadata=attn_metadata,
|
||||
)
|
||||
hidden_states = self.final_layer_norm.forward_native(hidden_states)
|
||||
hidden_states = self.final_layer_norm.forward(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
@@ -631,19 +631,13 @@ class UMT5EncoderModel(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", "q"),
|
||||
(".qkv_proj", ".k", "k"),
|
||||
(".qkv_proj", ".v", "v"),
|
||||
]
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: Set[str] = set()
|
||||
for name, loaded_weight in weights:
|
||||
loaded = False
|
||||
if "decoder" in name or "lm_head" in name:
|
||||
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)
|
||||
|
||||
@@ -10,17 +10,20 @@ from copy import deepcopy
|
||||
from typing import Any, Generator, Iterable, List, Optional, Tuple, cast
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn as nn
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
from transformers import AutoImageProcessor, AutoTokenizer
|
||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
|
||||
from fastvideo.v1.configs.models import EncoderConfig
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.hf_transformer_utils import get_diffusers_config
|
||||
from fastvideo.v1.models.loader.fsdp_load import maybe_load_fsdp_model
|
||||
from fastvideo.v1.models.loader.fsdp_load import (init_device_mesh,
|
||||
maybe_load_fsdp_model,
|
||||
shard_model)
|
||||
from fastvideo.v1.models.loader.utils import set_default_torch_dtype
|
||||
from fastvideo.v1.models.loader.weight_utils import (
|
||||
filter_duplicate_safetensors_files, filter_files_not_needed_for_inference,
|
||||
@@ -163,16 +166,19 @@ class TextEncoderLoader(ComponentLoader):
|
||||
return hf_folder, hf_weights_files, use_safetensors
|
||||
|
||||
def _get_weights_iterator(
|
||||
self, source: "Source"
|
||||
self,
|
||||
source: "Source",
|
||||
to_cpu: bool = True
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
"""Get an iterator for the model weights based on the load format."""
|
||||
hf_folder, hf_weights_files, use_safetensors = self._prepare_weights(
|
||||
source.model_or_path, source.fall_back_to_pt,
|
||||
source.allow_patterns_overrides)
|
||||
if use_safetensors:
|
||||
weights_iterator = safetensors_weights_iterator(hf_weights_files)
|
||||
weights_iterator = safetensors_weights_iterator(
|
||||
hf_weights_files, to_cpu)
|
||||
else:
|
||||
weights_iterator = pt_weights_iterator(hf_weights_files)
|
||||
weights_iterator = pt_weights_iterator(hf_weights_files, to_cpu)
|
||||
|
||||
if self.counter_before_loading_weights == 0.0:
|
||||
self.counter_before_loading_weights = time.perf_counter()
|
||||
@@ -181,10 +187,11 @@ class TextEncoderLoader(ComponentLoader):
|
||||
for (name, tensor) in weights_iterator)
|
||||
|
||||
def _get_all_weights(
|
||||
self,
|
||||
model_config: Any,
|
||||
model: nn.Module,
|
||||
model_path: str,
|
||||
self,
|
||||
model_config: Any,
|
||||
model: nn.Module,
|
||||
model_path: str,
|
||||
to_cpu: bool = True
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
primary_weights = TextEncoderLoader.Source(
|
||||
model_path,
|
||||
@@ -193,14 +200,14 @@ class TextEncoderLoader(ComponentLoader):
|
||||
allow_patterns_overrides=getattr(model, "allow_patterns_overrides",
|
||||
None),
|
||||
)
|
||||
yield from self._get_weights_iterator(primary_weights)
|
||||
yield from self._get_weights_iterator(primary_weights, to_cpu)
|
||||
|
||||
secondary_weights = cast(
|
||||
Iterable[TextEncoderLoader.Source],
|
||||
getattr(model, "secondary_weights", ()),
|
||||
)
|
||||
for source in secondary_weights:
|
||||
yield from self._get_weights_iterator(source)
|
||||
yield from self._get_weights_iterator(source, to_cpu)
|
||||
|
||||
def load(self, model_path: str, architecture: str,
|
||||
fastvideo_args: FastVideoArgs):
|
||||
@@ -233,16 +240,22 @@ class TextEncoderLoader(ComponentLoader):
|
||||
encoder_precision = fastvideo_args.pipeline_config.text_encoder_precisions[
|
||||
1]
|
||||
|
||||
target_device = get_torch_device()
|
||||
target_device = get_local_torch_device()
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(model_path, encoder_config, target_device,
|
||||
encoder_precision)
|
||||
fastvideo_args, encoder_precision)
|
||||
|
||||
def load_model(self,
|
||||
model_path: str,
|
||||
model_config: EncoderConfig,
|
||||
target_device: torch.device,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
dtype: str = "fp16"):
|
||||
use_cpu_offload = fastvideo_args.text_encoder_offload and len(
|
||||
getattr(model_config, "_fsdp_shard_conditions", [])) > 0
|
||||
|
||||
if fastvideo_args.text_encoder_offload:
|
||||
target_device = torch.device("cpu")
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
|
||||
with target_device:
|
||||
architectures = getattr(model_config, "architectures", [])
|
||||
@@ -251,12 +264,26 @@ class TextEncoderLoader(ComponentLoader):
|
||||
|
||||
weights_to_load = {name for name, _ in model.named_parameters()}
|
||||
loaded_weights = model.load_weights(
|
||||
self._get_all_weights(model_config, model, model_path))
|
||||
self._get_all_weights(model_config, model, model_path,
|
||||
use_cpu_offload))
|
||||
self.counter_after_loading_weights = time.perf_counter()
|
||||
logger.info(
|
||||
"Loading weights took %.2f seconds",
|
||||
self.counter_after_loading_weights -
|
||||
self.counter_before_loading_weights)
|
||||
|
||||
if use_cpu_offload:
|
||||
mesh = init_device_mesh(
|
||||
"cuda",
|
||||
mesh_shape=(1, dist.get_world_size()),
|
||||
mesh_dim_names=("offload", "replicate"),
|
||||
)
|
||||
shard_model(model,
|
||||
cpu_offload=True,
|
||||
reshard_after_forward=True,
|
||||
mesh=mesh["offload"],
|
||||
fsdp_shard_conditions=model._fsdp_shard_conditions,
|
||||
pin_cpu_memory=fastvideo_args.pin_cpu_memory)
|
||||
# We only enable strict check for non-quantized models
|
||||
# that have loaded weights tracking currently.
|
||||
# if loaded_weights is not None:
|
||||
@@ -290,10 +317,10 @@ class ImageEncoderLoader(TextEncoderLoader):
|
||||
encoder_config = fastvideo_args.pipeline_config.image_encoder_config
|
||||
encoder_config.update_model_arch(model_config)
|
||||
|
||||
target_device = get_torch_device()
|
||||
target_device = get_local_torch_device()
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(
|
||||
model_path, encoder_config, target_device,
|
||||
model_path, encoder_config, target_device, fastvideo_args,
|
||||
fastvideo_args.pipeline_config.image_encoder_precision)
|
||||
|
||||
|
||||
@@ -346,7 +373,7 @@ class VAELoader(ComponentLoader):
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]):
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
vae = vae_cls(vae_config).to(get_torch_device())
|
||||
vae = vae_cls(vae_config).to(get_local_torch_device())
|
||||
|
||||
# Find all safetensors files
|
||||
safetensors_list = glob.glob(
|
||||
@@ -405,7 +432,7 @@ class TransformerLoader(ComponentLoader):
|
||||
"hf_config": hf_config
|
||||
},
|
||||
weight_dir_list=safetensors_list,
|
||||
device=get_torch_device(),
|
||||
device=get_local_torch_device(),
|
||||
hsdp_replicate_dim=fastvideo_args.hsdp_replicate_dim,
|
||||
hsdp_shard_dim=fastvideo_args.hsdp_shard_dim,
|
||||
cpu_offload=fastvideo_args.use_cpu_offload,
|
||||
|
||||
@@ -69,6 +69,7 @@ def maybe_load_fsdp_model(
|
||||
fsdp_inference: bool = False,
|
||||
output_dtype: Optional[torch.dtype] = None,
|
||||
training_mode: bool = True,
|
||||
pin_cpu_memory: bool = True,
|
||||
) -> torch.nn.Module:
|
||||
"""
|
||||
Load the model with FSDP if is training, else load the model without FSDP.
|
||||
@@ -101,9 +102,12 @@ def maybe_load_fsdp_model(
|
||||
cpu_offload=cpu_offload,
|
||||
reshard_after_forward=True,
|
||||
mp_policy=mp_policy,
|
||||
mesh=device_mesh)
|
||||
mesh=device_mesh,
|
||||
fsdp_shard_conditions=model._fsdp_shard_conditions,
|
||||
pin_cpu_memory=pin_cpu_memory)
|
||||
|
||||
weight_iterator = safetensors_weights_iterator(weight_dir_list)
|
||||
weight_iterator = safetensors_weights_iterator(
|
||||
weight_dir_list, to_cpu=cpu_offload, async_broadcast=not cpu_offload)
|
||||
param_names_mapping_fn = get_param_names_mapping(model._param_names_mapping)
|
||||
load_model_from_full_model_state_dict(
|
||||
model,
|
||||
@@ -126,12 +130,13 @@ def maybe_load_fsdp_model(
|
||||
|
||||
def shard_model(
|
||||
model,
|
||||
*,
|
||||
cpu_offload: bool,
|
||||
reshard_after_forward: bool = True,
|
||||
mp_policy: Optional[MixedPrecisionPolicy] = None,
|
||||
dp_mesh: Optional[DeviceMesh] = None,
|
||||
mp_policy: Optional[MixedPrecisionPolicy] = MixedPrecisionPolicy(), # noqa
|
||||
mesh: Optional[DeviceMesh] = None,
|
||||
fsdp_shard_conditions: Optional[List[Callable[[str, nn.Module],
|
||||
bool]]] = None,
|
||||
pin_cpu_memory: bool = True,
|
||||
) -> None:
|
||||
"""
|
||||
Utility to shard a model with FSDP using the PyTorch Distributed fully_shard API.
|
||||
@@ -150,19 +155,28 @@ def shard_model(
|
||||
reshard_after_forward (bool): Whether to reshard parameters and buffers after
|
||||
the forward pass. Setting this to True corresponds to the FULL_SHARD sharding strategy
|
||||
from FSDP1, while setting it to False corresponds to the SHARD_GRAD_OP sharding strategy.
|
||||
dp_mesh (Optional[DeviceMesh]): Device mesh to use for FSDP sharding under multiple parallelism.
|
||||
mesh (Optional[DeviceMesh]): Device mesh to use for FSDP sharding under multiple parallelism.
|
||||
Default to None.
|
||||
fsdp_shard_conditions (Optional[List[Callable[[str, nn.Module], bool]]]): A list of functions to determine
|
||||
which modules to shard with FSDP.
|
||||
|
||||
Raises:
|
||||
ValueError: If no layer modules were sharded, indicating that no shard_condition was triggered.
|
||||
"""
|
||||
if fsdp_shard_conditions is None or len(fsdp_shard_conditions) == 0:
|
||||
logger.warning(
|
||||
"The FSDP shard condition list is empty or None. No modules will be sharded in %s",
|
||||
type(model).__name__)
|
||||
return
|
||||
|
||||
fsdp_kwargs = {
|
||||
"reshard_after_forward": reshard_after_forward,
|
||||
"mesh": mesh,
|
||||
"mp_policy": mp_policy,
|
||||
}
|
||||
if cpu_offload:
|
||||
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy()
|
||||
fsdp_kwargs["offload_policy"] = CPUOffloadPolicy(
|
||||
pin_memory=pin_cpu_memory)
|
||||
|
||||
# iterating in reverse to start with
|
||||
# lowest-level modules first
|
||||
@@ -172,7 +186,7 @@ def shard_model(
|
||||
for n, m in reversed(list(model.named_modules())):
|
||||
if any([
|
||||
shard_condition(n, m)
|
||||
for shard_condition in model._fsdp_shard_conditions
|
||||
for shard_condition in fsdp_shard_conditions
|
||||
]):
|
||||
fully_shard(m, **fsdp_kwargs)
|
||||
num_layers_sharded += 1
|
||||
@@ -181,7 +195,6 @@ def shard_model(
|
||||
raise ValueError(
|
||||
"No layer modules were sharded. Please check if shard conditions are working as expected."
|
||||
)
|
||||
|
||||
# Finally shard the entire model to account for any stragglers
|
||||
fully_shard(model, **fsdp_kwargs)
|
||||
|
||||
@@ -222,10 +235,17 @@ def load_model_from_full_model_state_dict(
|
||||
used_keys = set()
|
||||
sharded_sd = {}
|
||||
to_merge_params: DefaultDict[str, Dict[Any, Any]] = defaultdict(dict)
|
||||
reverse_param_names_mapping = {}
|
||||
assert param_names_mapping is not None
|
||||
|
||||
# iterate over all the weights to sync broadcast before use
|
||||
full_sd_iterator = list(full_sd_iterator) # type: ignore
|
||||
for source_param_name, full_tensor in full_sd_iterator:
|
||||
assert param_names_mapping is not None
|
||||
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
|
||||
source_param_name)
|
||||
reverse_param_names_mapping[target_param_name] = (source_param_name,
|
||||
merge_index,
|
||||
num_params_to_merge)
|
||||
used_keys.add(target_param_name)
|
||||
if merge_index is not None:
|
||||
to_merge_params[target_param_name][merge_index] = full_tensor
|
||||
@@ -260,6 +280,7 @@ def load_model_from_full_model_state_dict(
|
||||
sharded_tensor = sharded_tensor.cpu()
|
||||
sharded_sd[target_param_name] = nn.Parameter(sharded_tensor)
|
||||
|
||||
model._reverse_param_names_mapping = reverse_param_names_mapping
|
||||
unused_keys = set(meta_sd.keys()) - used_keys
|
||||
if unused_keys:
|
||||
logger.warning("Found new parameters in meta state dict: %s",
|
||||
|
||||
@@ -11,9 +11,11 @@ from typing import Generator, List, Optional, Tuple, Union
|
||||
import filelock
|
||||
import huggingface_hub.constants
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
from safetensors.torch import safe_open
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.v1.distributed.parallel_state import get_node_group
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -118,36 +120,77 @@ _BAR_FORMAT = "{desc}: {percentage:3.0f}% Completed | {n_fmt}/{total_fmt} [{elap
|
||||
|
||||
|
||||
def safetensors_weights_iterator(
|
||||
hf_weights_files: List[str]
|
||||
hf_weights_files: List[str],
|
||||
to_cpu: bool = False,
|
||||
async_broadcast: bool = False
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
"""Iterate over the weights in the model safetensor files."""
|
||||
enable_tqdm = not torch.distributed.is_initialized(
|
||||
) or torch.distributed.get_rank() == 0
|
||||
"""Iterate over the weights in the model safetensor files.
|
||||
Args:
|
||||
hf_weights_files: List of safetensor files to load.
|
||||
to_cpu: Whether to load the weights to CPU. If False, will load to the GPU device bound to the current process.
|
||||
async_broadcast: Whether to overlap loading from disk and broadcasting to other ranks. If True,
|
||||
must iterate over all the weights before use. Only use if to_cpu is False.
|
||||
"""
|
||||
local_rank = get_node_group().rank
|
||||
device = f"cuda:{local_rank}" if not to_cpu else "cpu"
|
||||
enable_tqdm = not torch.distributed.is_initialized() or get_node_group(
|
||||
).rank == 0
|
||||
assert not (async_broadcast
|
||||
and to_cpu), "Cannot broadcast weights when loading to CPU"
|
||||
|
||||
handles = []
|
||||
for st_file in tqdm(
|
||||
hf_weights_files,
|
||||
desc="Loading safetensors checkpoint shards",
|
||||
disable=not enable_tqdm,
|
||||
bar_format=_BAR_FORMAT,
|
||||
):
|
||||
with safe_open(st_file, framework="pt") as f:
|
||||
with safe_open(st_file, framework="pt", device=device) as f:
|
||||
for name in f.keys(): # noqa: SIM118
|
||||
param = f.get_tensor(name)
|
||||
if to_cpu:
|
||||
param = f.get_tensor(name)
|
||||
else:
|
||||
if local_rank == 0:
|
||||
param = f.get_tensor(name)
|
||||
else:
|
||||
shape = f.get_slice(name).get_shape()
|
||||
param = torch.empty(shape, device=device)
|
||||
# broadcast to local ranks
|
||||
# TODO(Wenxuan): scatter instead of broadcast
|
||||
if get_node_group().world_size > 1:
|
||||
group = get_node_group().device_group
|
||||
if async_broadcast:
|
||||
handle = dist.broadcast(param,
|
||||
src=dist.get_global_rank(
|
||||
group, 0),
|
||||
async_op=True)
|
||||
handles.append(handle)
|
||||
else:
|
||||
dist.broadcast(param,
|
||||
src=dist.get_global_rank(group, 0))
|
||||
yield name, param
|
||||
|
||||
if async_broadcast:
|
||||
for handle in handles:
|
||||
handle.wait()
|
||||
|
||||
|
||||
def pt_weights_iterator(
|
||||
hf_weights_files: List[str]
|
||||
hf_weights_files: List[str],
|
||||
to_cpu: bool = True # default to CPU for text encoder
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
"""Iterate over the weights in the model bin/pt files."""
|
||||
enable_tqdm = not torch.distributed.is_initialized(
|
||||
) or torch.distributed.get_rank() == 0
|
||||
local_rank = get_node_group().rank
|
||||
device = f"cuda:{local_rank}" if not to_cpu else "cpu"
|
||||
enable_tqdm = not torch.distributed.is_initialized() or get_node_group(
|
||||
).rank == 0
|
||||
for bin_file in tqdm(
|
||||
hf_weights_files,
|
||||
desc="Loading pt checkpoint shards",
|
||||
disable=not enable_tqdm,
|
||||
bar_format=_BAR_FORMAT,
|
||||
):
|
||||
state = torch.load(bin_file, map_location="cpu", weights_only=True)
|
||||
state = torch.load(bin_file, map_location=device, weights_only=True)
|
||||
yield from state.items()
|
||||
del state
|
||||
|
||||
|
||||
@@ -152,7 +152,6 @@ class TrainingBatch:
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None
|
||||
encoder_attention_mask: Optional[torch.Tensor] = None
|
||||
# i2v
|
||||
# extra_latents: Optional[Dict[str, Any]] = None
|
||||
preprocessed_image: Optional[torch.Tensor] = None
|
||||
image_embeds: Optional[torch.Tensor] = None
|
||||
image_latents: Optional[torch.Tensor] = None
|
||||
|
||||
@@ -18,7 +18,7 @@ from tqdm import tqdm
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.dataset import ValidationDataset, getdataset
|
||||
from fastvideo.v1.dataset.preprocessing_datasets import PreprocessBatch
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
@@ -104,7 +104,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
if strict:
|
||||
raise ValueError(
|
||||
f"Failed to convert tensor {tensor_name} to bytes: {e}"
|
||||
)
|
||||
) from e
|
||||
record[field] = b'' # Empty bytes for missing data
|
||||
else:
|
||||
record[field] = b'' # Empty bytes for missing data
|
||||
@@ -139,7 +139,8 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
except (ValueError, TypeError) as e:
|
||||
if strict:
|
||||
raise ValueError(
|
||||
f"Failed to convert field {field} to int: {e}")
|
||||
f"Failed to convert field {field} to int: {e}"
|
||||
) from e
|
||||
record[field] = 0
|
||||
else:
|
||||
record[field] = 0
|
||||
@@ -157,7 +158,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
if strict:
|
||||
raise ValueError(
|
||||
f"Failed to convert field {field} to float: {e}"
|
||||
)
|
||||
) from e
|
||||
record[field] = 0.0
|
||||
else:
|
||||
record[field] = 0.0
|
||||
@@ -211,8 +212,8 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
# Log unfilled fields as warning if not in strict mode
|
||||
if unfilled_fields:
|
||||
logger.warning(
|
||||
f"Some fields were not filled and got default values: {unfilled_fields}"
|
||||
)
|
||||
"Some fields were not filled and got default values: %s",
|
||||
unfilled_fields)
|
||||
|
||||
return record
|
||||
|
||||
@@ -221,7 +222,6 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
video_name: str,
|
||||
vae_latent: np.ndarray,
|
||||
text_embedding: np.ndarray,
|
||||
# text_attention_mask: np.ndarray,
|
||||
valid_data: Dict[str, Any],
|
||||
idx: int,
|
||||
extra_features: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||||
@@ -328,7 +328,8 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
# VAE
|
||||
with torch.autocast("cuda", dtype=torch.float32):
|
||||
latents = self.get_module("vae").encode(
|
||||
valid_data["pixel_values"].to(get_torch_device())).mean
|
||||
valid_data["pixel_values"].to(
|
||||
get_local_torch_device())).mean
|
||||
|
||||
# Get extra features if needed
|
||||
extra_features = self.get_extra_features(
|
||||
@@ -380,8 +381,6 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
# Convert tensors to numpy arrays
|
||||
vae_latent = latent.cpu().numpy()
|
||||
text_embedding = prompt_embeds[idx].cpu().numpy()
|
||||
# text_attention_mask = prompt_attention_mask[idx].cpu().numpy(
|
||||
# ).astype(np.uint8)
|
||||
|
||||
# Get extra features for this sample if needed
|
||||
sample_extra_features = {}
|
||||
@@ -398,7 +397,6 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
video_name=video_name,
|
||||
vae_latent=vae_latent,
|
||||
text_embedding=text_embedding,
|
||||
# text_attention_mask=text_attention_mask,
|
||||
valid_data=valid_data,
|
||||
idx=idx,
|
||||
extra_features=sample_extra_features)
|
||||
@@ -543,14 +541,13 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
valid_data["text"] = [prompt]
|
||||
|
||||
# Create record for Parquet dataset
|
||||
record = self.create_record(
|
||||
video_name=file_name,
|
||||
vae_latent=np.array([], dtype=np.float32),
|
||||
text_embedding=text_embedding,
|
||||
# text_attention_mask=text_attention_mask,
|
||||
valid_data=valid_data,
|
||||
idx=0,
|
||||
extra_features=sample_extra_features)
|
||||
record = self.create_record(video_name=file_name,
|
||||
vae_latent=np.array([],
|
||||
dtype=np.float32),
|
||||
text_embedding=text_embedding,
|
||||
valid_data=valid_data,
|
||||
idx=0,
|
||||
extra_features=sample_extra_features)
|
||||
batch_data.append(record)
|
||||
|
||||
logger.info("Saved validation sample: %s", file_name)
|
||||
|
||||
@@ -13,7 +13,7 @@ import torch
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_i2v
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.models.vision_utils import (get_default_height_width,
|
||||
@@ -58,7 +58,6 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
result_batch = self.image_encoding_stage(batch, fastvideo_args)
|
||||
clip_features = result_batch.image_embeds[0]
|
||||
|
||||
# image = self.pil_to_tensor(image)
|
||||
image = self.preprocess(
|
||||
image,
|
||||
vae_scale_factor=self.get_module("vae").spatial_compression_ratio,
|
||||
@@ -83,9 +82,8 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
|
||||
# TODO(will): move these to cpu at some point
|
||||
self.get_module("image_encoder").to(get_torch_device())
|
||||
# self.get_module("image_processor").to(get_torch_device())
|
||||
self.get_module("vae").to(get_torch_device())
|
||||
self.get_module("image_encoder").to(get_local_torch_device())
|
||||
self.get_module("vae").to(get_local_torch_device())
|
||||
|
||||
features = {}
|
||||
"""Get CLIP features from the first frame of each video."""
|
||||
@@ -109,7 +107,7 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
# Get CLIP features
|
||||
pixel_values = torch.cat(
|
||||
[img['pixel_values'] for img in processed_images],
|
||||
dim=0).to(get_torch_device())
|
||||
dim=0).to(get_local_torch_device())
|
||||
with torch.no_grad():
|
||||
image_inputs = {'pixel_values': pixel_values}
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
@@ -131,8 +129,8 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
height, width)
|
||||
],
|
||||
dim=2)
|
||||
video_condition = video_condition.to(device=get_torch_device(),
|
||||
dtype=torch.float32)
|
||||
video_condition = video_condition.to(
|
||||
device=get_local_torch_device(), dtype=torch.float32)
|
||||
video_conditions.append(video_condition)
|
||||
|
||||
video_conditions = torch.cat(video_conditions, dim=0)
|
||||
@@ -187,19 +185,16 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
video_name: str,
|
||||
vae_latent: np.ndarray,
|
||||
text_embedding: np.ndarray,
|
||||
# text_attention_mask: np.ndarray,
|
||||
valid_data: Optional[Dict[str, Any]],
|
||||
valid_data: Dict[str, Any],
|
||||
idx: int,
|
||||
extra_features: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||||
"""Create a record for the Parquet dataset with CLIP features."""
|
||||
record = super().create_record(
|
||||
video_name=video_name,
|
||||
vae_latent=vae_latent,
|
||||
text_embedding=text_embedding,
|
||||
# text_attention_mask=text_attention_mask,
|
||||
valid_data=valid_data,
|
||||
idx=idx,
|
||||
extra_features=extra_features)
|
||||
record = super().create_record(video_name=video_name,
|
||||
vae_latent=vae_latent,
|
||||
text_embedding=text_embedding,
|
||||
valid_data=valid_data,
|
||||
idx=idx,
|
||||
extra_features=extra_features)
|
||||
|
||||
if extra_features and "clip_feature" in extra_features:
|
||||
clip_feature = extra_features["clip_feature"]
|
||||
|
||||
@@ -59,12 +59,6 @@ if __name__ == "__main__":
|
||||
default=2,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--preprocess_text_batch_size",
|
||||
type=int,
|
||||
default=8,
|
||||
help="Batch size (per device) for the training dataloader.",
|
||||
)
|
||||
parser.add_argument("--samples_per_file", type=int, default=64)
|
||||
parser.add_argument("--flush_frequency",
|
||||
type=int,
|
||||
@@ -90,7 +84,7 @@ if __name__ == "__main__":
|
||||
type=str,
|
||||
default="google/t5-v1_1-xxl")
|
||||
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
|
||||
parser.add_argument("--cfg", type=float, default=0.0)
|
||||
parser.add_argument("--training_cfg_rate", type=float, default=0.0)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
|
||||
@@ -5,7 +5,7 @@ Decoding stage for diffusion pipelines.
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
@@ -61,7 +61,7 @@ class DecodingStage(PipelineStage):
|
||||
Returns:
|
||||
The batch with decoded outputs.
|
||||
"""
|
||||
self.vae = self.vae.to(get_torch_device())
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
|
||||
latents = batch.latents
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
|
||||
@@ -12,8 +12,9 @@ from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.v1.attention import get_attn_backend
|
||||
from fastvideo.v1.configs.pipelines.base import STA_Mode
|
||||
from fastvideo.v1.distributed import (get_sp_parallel_rank, get_sp_world_size,
|
||||
get_torch_device, get_world_group)
|
||||
from fastvideo.v1.distributed import (get_local_torch_device,
|
||||
get_sp_parallel_rank, get_sp_world_size,
|
||||
get_world_group)
|
||||
from fastvideo.v1.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
@@ -192,7 +193,7 @@ class DenoisingStage(PipelineStage):
|
||||
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
|
||||
latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=get_torch_device(),
|
||||
device=get_local_torch_device(),
|
||||
).to(target_dtype) *
|
||||
1000.0 if fastvideo_args.pipeline_config.embedded_cfg_scale
|
||||
is not None else None)
|
||||
|
||||
@@ -7,7 +7,7 @@ from typing import Optional
|
||||
import PIL.Image
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
@@ -49,7 +49,7 @@ class EncodingStage(PipelineStage):
|
||||
Returns:
|
||||
The batch with encoded outputs.
|
||||
"""
|
||||
self.vae = self.vae.to(get_torch_device())
|
||||
self.vae = self.vae.to(get_local_torch_device())
|
||||
|
||||
assert batch.height is not None
|
||||
assert batch.width is not None
|
||||
@@ -65,10 +65,12 @@ class EncodingStage(PipelineStage):
|
||||
image,
|
||||
vae_scale_factor=self.vae.spatial_compression_ratio,
|
||||
height=batch.height,
|
||||
width=batch.width).to(get_torch_device(), dtype=torch.float32)
|
||||
width=batch.width).to(get_local_torch_device(),
|
||||
dtype=torch.float32)
|
||||
|
||||
image = image.unsqueeze(2)
|
||||
else:
|
||||
# assumes image is loaded from parquet file and used for validation
|
||||
image = image.transpose(1, 2)
|
||||
logger.info("image: %s", image.shape)
|
||||
video_condition = torch.cat([
|
||||
@@ -77,7 +79,7 @@ class EncodingStage(PipelineStage):
|
||||
batch.num_frames - 1, batch.height, batch.width)
|
||||
],
|
||||
dim=2)
|
||||
video_condition = video_condition.to(device=get_torch_device(),
|
||||
video_condition = video_condition.to(device=get_local_torch_device(),
|
||||
dtype=torch.float32)
|
||||
|
||||
# Setup VAE precision
|
||||
@@ -101,6 +103,7 @@ class EncodingStage(PipelineStage):
|
||||
generator = batch.generator
|
||||
if generator is None:
|
||||
raise ValueError("Generator must be provided")
|
||||
# latent_condition = self.retrieve_latents(encoder_output, generator, sample_mode="argmax")
|
||||
latent_condition = self.retrieve_latents(encoder_output, generator)
|
||||
|
||||
# Apply shifting if needed
|
||||
|
||||
@@ -7,7 +7,7 @@ This module contains implementations of image encoding stages for diffusion pipe
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -55,12 +55,12 @@ class ImageEncodingStage(PipelineStage):
|
||||
The batch with encoded prompt embeddings.
|
||||
"""
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.image_encoder = self.image_encoder.to(get_torch_device())
|
||||
self.image_encoder = self.image_encoder.to(get_local_torch_device())
|
||||
|
||||
image = batch.pil_image
|
||||
|
||||
image_inputs = self.image_processor(
|
||||
images=image, return_tensors="pt").to(get_torch_device())
|
||||
images=image, return_tensors="pt").to(get_local_torch_device())
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs = self.image_encoder(**image_inputs)
|
||||
image_embeds = outputs.last_hidden_state
|
||||
|
||||
@@ -5,7 +5,7 @@ Latent preparation stage for diffusion pipelines.
|
||||
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
@@ -62,7 +62,7 @@ class LatentPreparationStage(PipelineStage):
|
||||
|
||||
# Get required parameters
|
||||
dtype = batch.prompt_embeds[0].dtype
|
||||
device = get_torch_device()
|
||||
device = get_local_torch_device()
|
||||
generator = batch.generator
|
||||
latents = batch.latents
|
||||
num_frames = latent_num_frames if latent_num_frames is not None else batch.num_frames
|
||||
|
||||
@@ -5,9 +5,7 @@ Prompt encoding stages for diffusion pipelines.
|
||||
This module contains implementations of prompt encoding stages for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
@@ -62,8 +60,6 @@ class TextEncodingStage(PipelineStage):
|
||||
fastvideo_args.pipeline_config.text_encoder_configs,
|
||||
fastvideo_args.pipeline_config.preprocess_text_funcs,
|
||||
fastvideo_args.pipeline_config.postprocess_text_funcs):
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
text_encoder = text_encoder.to(get_torch_device())
|
||||
|
||||
assert isinstance(batch.prompt, (str, list))
|
||||
if isinstance(batch.prompt, str):
|
||||
@@ -71,8 +67,9 @@ class TextEncodingStage(PipelineStage):
|
||||
texts = []
|
||||
for prompt_str in batch.prompt:
|
||||
texts.append(preprocess_func(prompt_str))
|
||||
text_inputs = tokenizer(
|
||||
texts, **encoder_config.tokenizer_kwargs).to(get_torch_device())
|
||||
text_inputs = tokenizer(texts,
|
||||
**encoder_config.tokenizer_kwargs).to(
|
||||
get_local_torch_device())
|
||||
input_ids = text_inputs["input_ids"]
|
||||
attention_mask = text_inputs["attention_mask"]
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
@@ -91,8 +88,8 @@ class TextEncodingStage(PipelineStage):
|
||||
assert isinstance(batch.negative_prompt, str)
|
||||
negative_text = preprocess_func(batch.negative_prompt)
|
||||
negative_text_inputs = tokenizer(
|
||||
negative_text,
|
||||
**encoder_config.tokenizer_kwargs).to(get_torch_device())
|
||||
negative_text, **encoder_config.tokenizer_kwargs).to(
|
||||
get_local_torch_device())
|
||||
negative_input_ids = negative_text_inputs["input_ids"]
|
||||
negative_attention_mask = negative_text_inputs["attention_mask"]
|
||||
with set_forward_context(current_timestep=0,
|
||||
@@ -110,9 +107,8 @@ class TextEncodingStage(PipelineStage):
|
||||
batch.negative_attention_mask.append(
|
||||
negative_attention_mask)
|
||||
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
if fastvideo_args.text_encoder_offload:
|
||||
text_encoder.to('cpu')
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ This module contains implementations of timestep preparation stages for diffusio
|
||||
|
||||
import inspect
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
@@ -45,7 +45,7 @@ class TimestepPreparationStage(PipelineStage):
|
||||
The batch with prepared timesteps.
|
||||
"""
|
||||
scheduler = self.scheduler
|
||||
device = get_torch_device()
|
||||
device = get_local_torch_device()
|
||||
num_inference_steps = batch.num_inference_steps
|
||||
timesteps = batch.timesteps
|
||||
sigmas = batch.sigmas
|
||||
|
||||
@@ -14,7 +14,7 @@ from typing import Any, Dict
|
||||
import torch
|
||||
from huggingface_hub import hf_hub_download
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.distributed import get_local_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.encoders.bert import HunyuanClip # type: ignore
|
||||
@@ -78,7 +78,7 @@ class StepVideoPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""
|
||||
Initialize the pipeline.
|
||||
"""
|
||||
target_device = get_torch_device()
|
||||
target_device = get_local_torch_device()
|
||||
llm_dir = os.path.join(self.model_path, "step_llm")
|
||||
clip_dir = os.path.join(self.model_path, "hunyuan_clip")
|
||||
text_enc = self.build_llm(llm_dir, target_device)
|
||||
|
||||
@@ -6,7 +6,7 @@ import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from transformers import AutoConfig
|
||||
|
||||
import gc
|
||||
from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
|
||||
load_tokenizer)
|
||||
# from fastvideo.v1.models.hunyuan.text_encoder import load_text_encoder, load_tokenizer
|
||||
@@ -16,6 +16,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
from fastvideo.v1.configs.models.encoders import CLIPTextConfig
|
||||
from torch.distributed.tensor import DTensor
|
||||
from torch.testing import assert_close
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -66,7 +68,6 @@ def test_clip_encoder():
|
||||
# Load the HuggingFace implementation directly
|
||||
# model2 = CLIPTextModel(hf_config)
|
||||
# model2 = model2.to(torch.float16)
|
||||
model2 = model2.to(device)
|
||||
model2.eval()
|
||||
|
||||
# Sanity check weights between the two models
|
||||
@@ -78,19 +79,20 @@ def test_clip_encoder():
|
||||
logger.info("Model1 has %d parameters", len(params1))
|
||||
logger.info("Model2 has %d parameters", len(params2))
|
||||
|
||||
# Compare a few key parameters
|
||||
|
||||
# weight_diffs = []
|
||||
# for (name1, param1), (name2, param2) in zip(
|
||||
# sorted(params1.items()), sorted(params2.items())
|
||||
# ):
|
||||
# # if len(weight_diffs) < 5: # Just check a few parameters
|
||||
# max_diff = torch.max(torch.abs(param1 - param2)).item()
|
||||
# mean_diff = torch.mean(torch.abs(param1 - param2)).item()
|
||||
# weight_diffs.append((name1, name2, max_diff, mean_diff))
|
||||
# logger.info(f"Parameter: {name1} vs {name2}")
|
||||
# logger.info(f" Max diff: {max_diff}, Mean diff: {mean_diff}")
|
||||
|
||||
for name1, param1 in sorted(params1.items()):
|
||||
name2 = name1
|
||||
skip = False
|
||||
for param_name, weight_name, shard_id in model2.config.arch_config.stacked_params_mapping:
|
||||
if weight_name not in name1:
|
||||
skip = True
|
||||
# stacked params are more troublesome
|
||||
if skip:
|
||||
continue
|
||||
param2 = params2[name2]
|
||||
param2 = param2.to_local().to(device) if isinstance(param2, DTensor) else param2.to(device)
|
||||
assert_close(param1, param2, atol=1e-4, rtol=1e-4)
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
# Load tokenizer
|
||||
tokenizer, _ = load_tokenizer(tokenizer_type="clipL",
|
||||
tokenizer_path=args.model_path,
|
||||
@@ -168,5 +170,5 @@ def test_clip_encoder():
|
||||
f"Pooler outputs differ significantly: mean diff = {mean_diff_pooler.item()}"
|
||||
assert max_diff_hidden < 1e-1, \
|
||||
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
|
||||
assert max_diff_pooler < 1e-2, \
|
||||
assert max_diff_pooler < 2e-2, \
|
||||
f"Pooler outputs differ significantly: max diff = {max_diff_pooler.item()}"
|
||||
|
||||
@@ -5,7 +5,7 @@ import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from transformers import AutoConfig
|
||||
|
||||
import gc
|
||||
from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
|
||||
load_tokenizer)
|
||||
from fastvideo.v1.configs.pipelines import PipelineConfig
|
||||
@@ -15,7 +15,8 @@ from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
from fastvideo.v1.configs.models.encoders import LlamaConfig
|
||||
|
||||
from torch.distributed.tensor import DTensor
|
||||
from torch.testing import assert_close
|
||||
logger = init_logger(__name__)
|
||||
|
||||
os.environ["MASTER_ADDR"] = "localhost"
|
||||
@@ -62,7 +63,6 @@ def test_llama_encoder():
|
||||
|
||||
# Convert to float16 and move to device
|
||||
# model2 = model2.to(torch.float16)
|
||||
model2 = model2.to(device)
|
||||
model2.eval()
|
||||
|
||||
# Sanity check weights between the two models
|
||||
@@ -77,34 +77,28 @@ def test_llama_encoder():
|
||||
# Compare a few key parameters
|
||||
weight_diffs = []
|
||||
# check if embed_tokens are the same
|
||||
print(model1.embed_tokens.weight.shape, model2.embed_tokens.weight.shape)
|
||||
device = model1.embed_tokens.weight.device
|
||||
assert torch.allclose(model1.embed_tokens.weight,
|
||||
model2.embed_tokens.weight)
|
||||
model2.embed_tokens.weight.to_local().to(device) if isinstance(model2.embed_tokens.weight, DTensor) else model2.embed_tokens.weight.to(device))
|
||||
weights = [
|
||||
"layers.{}.input_layernorm.weight",
|
||||
"layers.{}.post_attention_layernorm.weight"
|
||||
]
|
||||
# for (name1, param1), (name2, param2) in zip(
|
||||
# sorted(params1.items()), sorted(params2.items())
|
||||
# ):
|
||||
for layer_idx in range(hf_config.num_hidden_layers):
|
||||
for w in weights:
|
||||
name1 = w.format(layer_idx)
|
||||
name2 = w.format(layer_idx)
|
||||
p1 = params1[name1]
|
||||
p2 = params2[name2]
|
||||
# print(type(p2))
|
||||
if "gate_up" in name2:
|
||||
# print("skipping gate_up")
|
||||
continue
|
||||
try:
|
||||
# logger.info(f"Parameter: {name1} vs {name2}")
|
||||
max_diff = torch.max(torch.abs(p1 - p2)).item()
|
||||
mean_diff = torch.mean(torch.abs(p1 - p2)).item()
|
||||
weight_diffs.append((name1, name2, max_diff, mean_diff))
|
||||
# logger.info(f" Max diff: {max_diff}, Mean diff: {mean_diff}")
|
||||
except Exception as e:
|
||||
logger.info("Error comparing %s and %s: %s", name1, name2, e)
|
||||
|
||||
for name1, param1 in sorted(params1.items()):
|
||||
name2 = name1
|
||||
skip = False
|
||||
for param_name, weight_name, shard_id in model2.config.arch_config.stacked_params_mapping:
|
||||
if weight_name not in name1:
|
||||
skip = True
|
||||
# stacked params are more troublesome
|
||||
if skip:
|
||||
continue
|
||||
param2 = params2[name2]
|
||||
param2 = param2.to_local().to(device) if isinstance(param2, DTensor) else param2.to(device)
|
||||
assert_close(param1, param2, atol=1e-4, rtol=1e-4)
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
tokenizer, _ = load_tokenizer(tokenizer_type="llm",
|
||||
tokenizer_path=TOKENIZER_PATH,
|
||||
|
||||
@@ -4,6 +4,8 @@ import os
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
from torch.distributed.tensor import DTensor
|
||||
from torch.testing import assert_close
|
||||
from transformers import AutoConfig, AutoTokenizer, UMT5EncoderModel
|
||||
|
||||
from fastvideo.v1.configs.pipelines import PipelineConfig
|
||||
@@ -41,13 +43,13 @@ def test_t5_encoder():
|
||||
tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
|
||||
|
||||
|
||||
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH, pipeline_config=PipelineConfig(text_encoder_configs=(T5Config(),), text_encoder_precisions=(precision_str,)))
|
||||
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH,
|
||||
pipeline_config=PipelineConfig(text_encoder_configs=(T5Config(),),
|
||||
text_encoder_precisions=(precision_str,)),
|
||||
pin_cpu_memory=False)
|
||||
loader = TextEncoderLoader()
|
||||
model2 = loader.load(TEXT_ENCODER_PATH, "", args)
|
||||
|
||||
# Convert to float16 and move to device
|
||||
# model2 = model2.to(precision)
|
||||
model2 = model2.to(device)
|
||||
model2 = model2.to(precision)
|
||||
model2.eval()
|
||||
|
||||
# Sanity check weights between the two models
|
||||
@@ -64,23 +66,17 @@ def test_t5_encoder():
|
||||
weights = ["encoder.block.{}.layer.0.layer_norm.weight", "encoder.block.{}.layer.0.SelfAttention.relative_attention_bias.weight", \
|
||||
"encoder.block.{}.layer.0.SelfAttention.o.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_0.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_1.weight",\
|
||||
"encoder.block.{}.layer.1.DenseReluDense.wo.weight", \
|
||||
"encoder.block.{}.layer.1.layer_norm.weight", "encoder.final_layer_norm.weight", "shared.weight"]
|
||||
"encoder.block.{}.layer.1.layer_norm.weight", "encoder.final_layer_norm.weight"]
|
||||
|
||||
for idx in range(hf_config.num_hidden_layers):
|
||||
for w in weights:
|
||||
name1 = w.format(idx)
|
||||
name2 = w.format(idx)
|
||||
p1 = params1[name1]
|
||||
p2 = params2[name2]
|
||||
assert p1.dtype == p2.dtype
|
||||
try:
|
||||
logger.info("Parameter: %s vs %s", name1, name2)
|
||||
max_diff = torch.max(torch.abs(p1 - p2)).item()
|
||||
mean_diff = torch.mean(torch.abs(p1 - p2)).item()
|
||||
weight_diffs.append((name1, name2, max_diff, mean_diff))
|
||||
logger.info(" Max diff: %s, Mean diff: %s", max_diff,
|
||||
mean_diff)
|
||||
except Exception as e:
|
||||
logger.info("Error comparing %s and %s: %s", name1, name2, e)
|
||||
p2 = (p2.to_local() if isinstance(p2, DTensor) else p2).to(p1)
|
||||
assert_close(p1, p2, atol=1e-4, rtol=1e-4)
|
||||
|
||||
|
||||
# Test with some sample prompts
|
||||
prompts = [
|
||||
|
||||
@@ -4,95 +4,83 @@ app = modal.App()
|
||||
|
||||
import os
|
||||
|
||||
image_version = os.getenv("IMAGE_VERSION", "latest")
|
||||
image_version = os.getenv("IMAGE_VERSION")
|
||||
image_tag = f"ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:{image_version}"
|
||||
print(f"Using image: {image_tag}")
|
||||
|
||||
image = (
|
||||
modal.Image.from_registry(image_tag, add_python="3.12")
|
||||
.run_commands("rm -rf /FastVideo")
|
||||
.apt_install("cmake", "pkg-config", "build-essential", "curl", "libssl-dev")
|
||||
.run_commands("curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y --default-toolchain stable")
|
||||
.run_commands("echo 'source ~/.cargo/env' >> ~/.bashrc")
|
||||
.env({"PATH": "/root/.cargo/bin:$PATH"})
|
||||
.run_commands("/bin/bash -c 'source $HOME/.local/bin/env && source /opt/venv/bin/activate && cd /FastVideo && uv pip install -e .[test]'")
|
||||
.env({
|
||||
"PATH": "/root/.cargo/bin:$PATH",
|
||||
"BUILDKITE_REPO": os.environ.get("BUILDKITE_REPO", ""),
|
||||
"BUILDKITE_COMMIT": os.environ.get("BUILDKITE_COMMIT", ""),
|
||||
})
|
||||
)
|
||||
|
||||
def run_test(pytest_command: str):
|
||||
"""Helper function to run a test suite with custom pytest command"""
|
||||
import subprocess
|
||||
import sys
|
||||
import os
|
||||
|
||||
git_repo = os.environ.get("BUILDKITE_REPO")
|
||||
git_commit = os.environ.get("BUILDKITE_COMMIT")
|
||||
|
||||
print(f"Cloning repository: {git_repo}")
|
||||
print(f"Checking out commit: {git_commit}")
|
||||
|
||||
command = f"""
|
||||
source $HOME/.local/bin/env &&
|
||||
source /opt/venv/bin/activate &&
|
||||
git clone {git_repo} /FastVideo &&
|
||||
cd /FastVideo &&
|
||||
git checkout {git_commit} &&
|
||||
uv pip install -e .[test] &&
|
||||
{pytest_command}
|
||||
"""
|
||||
|
||||
result = subprocess.run([
|
||||
"/bin/bash", "-c", command
|
||||
], stdout=sys.stdout, stderr=sys.stderr, check=False)
|
||||
|
||||
sys.exit(result.returncode)
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1800)
|
||||
def run_encoder_tests():
|
||||
"""Run encoder tests on L40S GPU"""
|
||||
import subprocess
|
||||
import sys
|
||||
import os
|
||||
|
||||
os.chdir("/FastVideo")
|
||||
|
||||
command = """
|
||||
source /opt/venv/bin/activate &&
|
||||
pytest ./fastvideo/v1/tests/encoders -s
|
||||
"""
|
||||
|
||||
result = subprocess.run([
|
||||
"/bin/bash", "-c", command
|
||||
], stdout=sys.stdout, stderr=sys.stderr, check=False)
|
||||
|
||||
sys.exit(result.returncode)
|
||||
run_test("pytest ./fastvideo/v1/tests/encoders -vs")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1800)
|
||||
def run_vae_tests():
|
||||
"""Run VAE tests on L40S GPU"""
|
||||
import subprocess
|
||||
import sys
|
||||
import os
|
||||
|
||||
os.chdir("/FastVideo")
|
||||
|
||||
command = """
|
||||
source /opt/venv/bin/activate &&
|
||||
pytest ./fastvideo/v1/tests/vaes -s
|
||||
"""
|
||||
|
||||
result = subprocess.run([
|
||||
"/bin/bash", "-c", command
|
||||
], stdout=sys.stdout, stderr=sys.stderr, check=False)
|
||||
|
||||
sys.exit(result.returncode)
|
||||
run_test("pytest ./fastvideo/v1/tests/vaes -vs")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1800)
|
||||
def run_transformer_tests():
|
||||
"""Run transformer tests on L40S GPU"""
|
||||
import subprocess
|
||||
import sys
|
||||
import os
|
||||
|
||||
os.chdir("/FastVideo")
|
||||
|
||||
command = """
|
||||
source /opt/venv/bin/activate &&
|
||||
pytest ./fastvideo/v1/tests/transformers -s
|
||||
"""
|
||||
|
||||
result = subprocess.run([
|
||||
"/bin/bash", "-c", command
|
||||
], stdout=sys.stdout, stderr=sys.stderr, check=False)
|
||||
|
||||
sys.exit(result.returncode)
|
||||
run_test("pytest ./fastvideo/v1/tests/transformers -vs")
|
||||
|
||||
@app.function(gpu="L40S:2", image=image, timeout=3600)
|
||||
def run_ssim_tests():
|
||||
"""Run SSIM tests on 2x L40S GPUs"""
|
||||
import subprocess
|
||||
import sys
|
||||
import os
|
||||
|
||||
os.chdir("/FastVideo")
|
||||
|
||||
command = """
|
||||
source /opt/venv/bin/activate &&
|
||||
pytest ./fastvideo/v1/tests/ssim -vs
|
||||
"""
|
||||
|
||||
result = subprocess.run([
|
||||
"/bin/bash", "-c", command
|
||||
], stdout=sys.stdout, stderr=sys.stderr, check=False)
|
||||
|
||||
sys.exit(result.returncode)
|
||||
run_test("pytest ./fastvideo/v1/tests/ssim -vs")
|
||||
|
||||
@app.function(gpu="L40S:4", image=image, timeout=1800, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
def run_training_tests():
|
||||
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/v1/tests/training/Vanilla -srP")
|
||||
|
||||
@app.function(gpu="H100:1", image=image, timeout=1800, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
|
||||
def run_training_tests_VSA():
|
||||
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/v1/tests/training/VSA -srP")
|
||||
|
||||
@app.function(gpu="H100:1", image=image, timeout=1800)
|
||||
def run_inference_tests_STA():
|
||||
run_test("pytest ./fastvideo/v1/tests/inference/STA -srP")
|
||||
|
||||
@app.function(gpu="H100:1", image=image, timeout=1800)
|
||||
def run_precision_tests_STA():
|
||||
run_test("python csrc/attn/tests/test_sta.py")
|
||||
|
||||
@app.function(gpu="H100:1", image=image, timeout=1800)
|
||||
def run_precision_tests_VSA():
|
||||
run_test("python csrc/attn/tests/test_block_sparse.py")
|
||||
|
||||
@@ -122,7 +122,7 @@ def run_training():
|
||||
"--checkpoints_total_limit", "3",
|
||||
"--allow_tf32",
|
||||
"--ema_start_step", "0",
|
||||
"--cfg", "0.0",
|
||||
"--training_cfg_rate", "0.1",
|
||||
"--output_dir", LOCAL_OUTPUT_DIR,
|
||||
"--tracker_project_name", "wan_i2v_finetune_overfit_ci",
|
||||
"--num_height", "480",
|
||||
|
||||
@@ -31,12 +31,9 @@ LOCAL_OUTPUT_DIR = Path(os.path.join(DATA_DIR, "outputs"))
|
||||
def download_data():
|
||||
# create the data dir if it doesn't exist
|
||||
data_dir = Path(DATA_DIR)
|
||||
if data_dir.exists():
|
||||
print(f"Removing existing data directory at {data_dir}")
|
||||
shutil.rmtree(data_dir)
|
||||
|
||||
print(f"Creating data directory at {data_dir}")
|
||||
os.makedirs(data_dir)
|
||||
os.makedirs(data_dir, exist_ok=True)
|
||||
|
||||
print(f"Downloading raw dataset to {LOCAL_RAW_DATA_DIR}...")
|
||||
try:
|
||||
@@ -122,7 +119,7 @@ def run_training():
|
||||
"--checkpoints_total_limit", "3",
|
||||
"--allow_tf32",
|
||||
"--ema_start_step", "0",
|
||||
"--cfg", "0.0",
|
||||
"--training_cfg_rate", "0.0",
|
||||
"--output_dir", LOCAL_OUTPUT_DIR,
|
||||
"--tracker_project_name", "wan_finetune_overfit_ci",
|
||||
"--num_height", "480",
|
||||
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -1,13 +1,14 @@
|
||||
The reference videos in the `reference_videos` directory are used as part of an e2e test to ensure consistency in video generation quality across code changes. `test_inference_similarity.py` compares newly generated videos against these references using Structural Similarity Index (SSIM) metrics to detect any regressions in visual quality across code changes.
|
||||
The reference videos in the `*_reference_videos` directory are used as part of an e2e test to ensure consistency in video generation quality across code changes. `test_inference_similarity.py` compares newly generated videos against these references using Structural Similarity Index (SSIM) metrics to detect any regressions in visual quality across code changes.
|
||||
|
||||
`reference_videos/FastHunyuan-diffusers/FLASH_ATTN/` videos were generated on commit `66107fd5b8469fed25972feb632cd48887dac451`.
|
||||
`reference_videos/FastHunyuan-diffusers/TORCH_SDPA/` videos were generated on commit `4ea008b8a16d7f5678a44b187ebdd7d9d0416ff1`.
|
||||
`reference_videos/Wan2.1-T2V-1.3B-Diffusers` videos were generated on commit `d085770a70988c7b26632a0c3123c24a57f7ca77`.
|
||||
`reference_videos/Wan2.1-I2V-14B-480P-Diffusers` videos were generated on commit `d085770a70988c7b26632a0c3123c24a57f7ca77`.
|
||||
`A40_reference_videos` are generated on A40s and so on.
|
||||
|
||||
run `bash update_reference_videos.sh` from inside the `fastvideo/v1/tests/ssim/` directory after running `test_inference_similarity.py` to update reference videos. Note: make sure to update the path to the corresponding device.
|
||||
|
||||
all reference videos are were generated on commit `4aeabbc629e0edf91477e80e795e7bb1823c71cb`
|
||||
|
||||
## Generation Details
|
||||
|
||||
2 x NVIDIA A40 GPUs
|
||||
2 x NVIDIA L40S GPUs
|
||||
|
||||
## Generation Parameters
|
||||
|
||||
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -2,6 +2,7 @@
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
import pytest
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
@@ -11,6 +12,14 @@ from fastvideo.v1.worker.multiproc_executor import MultiprocExecutor
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
device_name = torch.cuda.get_device_name()
|
||||
device_reference_folder_suffix = '_reference_videos'
|
||||
|
||||
if "A40" in device_name:
|
||||
device_reference_folder = "A40" + device_reference_folder_suffix
|
||||
elif "L40S" in device_name:
|
||||
device_reference_folder = "L40S" + device_reference_folder_suffix
|
||||
|
||||
# Base parameters from the shell script
|
||||
HUNYUAN_PARAMS = {
|
||||
"num_gpus": 2,
|
||||
@@ -188,8 +197,8 @@ def test_i2v_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
assert os.path.exists(
|
||||
output_dir), f"Output video was not generated at {output_dir}"
|
||||
|
||||
reference_folder = os.path.join(script_dir, 'reference_videos', model_id, ATTENTION_BACKEND)
|
||||
|
||||
reference_folder = os.path.join(script_dir, device_reference_folder, model_id, ATTENTION_BACKEND)
|
||||
|
||||
if not os.path.exists(reference_folder):
|
||||
logger.error("Reference folder missing")
|
||||
raise FileNotFoundError(
|
||||
@@ -288,8 +297,8 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
assert os.path.exists(
|
||||
output_dir), f"Output video was not generated at {output_dir}"
|
||||
|
||||
reference_folder = os.path.join(script_dir, 'reference_videos', model_id, ATTENTION_BACKEND)
|
||||
|
||||
reference_folder = os.path.join(script_dir, device_reference_folder, model_id, ATTENTION_BACKEND)
|
||||
|
||||
if not os.path.exists(reference_folder):
|
||||
logger.error("Reference folder missing")
|
||||
raise FileNotFoundError(
|
||||
|
||||
@@ -0,0 +1,63 @@
|
||||
#!/bin/bash
|
||||
|
||||
# Script to update reference videos using videos from generated_videos directory
|
||||
# Both directories should exist in the same directory as this script
|
||||
|
||||
set -e # Exit on any error
|
||||
|
||||
# Define directory paths
|
||||
GENERATED_DIR="generated_videos"
|
||||
REFERENCE_DIR="set_me_to_correct_path"
|
||||
|
||||
# Colors for output
|
||||
RED='\033[0;31m'
|
||||
GREEN='\033[0;32m'
|
||||
YELLOW='\033[1;33m'
|
||||
NC='\033[0m' # No Color
|
||||
|
||||
echo -e "${YELLOW}Starting reference video update...${NC}"
|
||||
|
||||
# Check if generated_videos directory exists
|
||||
if [ ! -d "$GENERATED_DIR" ]; then
|
||||
echo -e "${RED}Error: $GENERATED_DIR directory not found!${NC}"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Check if reference_videos directory exists
|
||||
if [ ! -d "$REFERENCE_DIR" ]; then
|
||||
echo -e "${RED}Error: $REFERENCE_DIR directory not found!${NC}"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
# Function to copy videos recursively
|
||||
copy_videos() {
|
||||
local src_dir="$1"
|
||||
local dst_dir="$2"
|
||||
|
||||
# Find all video files in the source directory
|
||||
find "$src_dir" -type f \( -name "*.mp4" -o -name "*.avi" -o -name "*.mov" -o -name "*.mkv" -o -name "*.webm" -o -name "*.flv" \) | while read -r video_file; do
|
||||
# Get relative path from source directory
|
||||
relative_path="${video_file#$src_dir/}"
|
||||
|
||||
# Construct destination path
|
||||
dst_file="$dst_dir/$relative_path"
|
||||
|
||||
# Create destination directory if it doesn't exist
|
||||
dst_file_dir=$(dirname "$dst_file")
|
||||
mkdir -p "$dst_file_dir"
|
||||
|
||||
# Copy the video file
|
||||
echo -e "${GREEN}Copying: $relative_path${NC}"
|
||||
cp "$video_file" "$dst_file"
|
||||
done
|
||||
}
|
||||
|
||||
# Perform the copy operation
|
||||
echo -e "${YELLOW}Copying videos from $GENERATED_DIR to $REFERENCE_DIR...${NC}"
|
||||
copy_videos "$GENERATED_DIR" "$REFERENCE_DIR"
|
||||
|
||||
echo -e "${GREEN}Reference videos updated successfully!${NC}"
|
||||
|
||||
# Show summary
|
||||
video_count=$(find "$GENERATED_DIR" -type f \( -name "*.mp4" -o -name "*.avi" -o -name "*.mov" -o -name "*.mkv" -o -name "*.webm" -o -name "*.flv" \) | wc -l)
|
||||
echo -e "${YELLOW}Total videos processed: $video_count${NC}"
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user