Compare commits
44
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ad16289871 | ||
|
|
2a41da1e6b | ||
|
|
32133171da | ||
|
|
19674c6f29 | ||
|
|
508afb7002 | ||
|
|
288ea88105 | ||
|
|
eb0f1318f3 | ||
|
|
ce9b5910cc | ||
|
|
d0e5a6214a | ||
|
|
834562b2db | ||
|
|
060cc7b9ba | ||
|
|
6c58a5ba62 | ||
|
|
48d9f61f86 | ||
|
|
5f938b5844 | ||
|
|
74da2a7370 | ||
|
|
580d6dfe1f | ||
|
|
344e43006a | ||
|
|
c5155b256e | ||
|
|
e005c7f3ac | ||
|
|
ff5a79ef60 | ||
|
|
ab01dc4ba5 | ||
|
|
285a950c1b | ||
|
|
46a0a85d85 | ||
|
|
4aeabbc629 | ||
|
|
949bb5c835 | ||
|
|
aab74c1271 | ||
|
|
f89d86944f | ||
|
|
8741d204a5 | ||
|
|
cdc85f58a8 | ||
|
|
0262d2f089 | ||
|
|
62c0343465 | ||
|
|
1e1a023fb0 | ||
|
|
1d2517ad8e | ||
|
|
d41186cb4a | ||
|
|
78e0c7eec9 | ||
|
|
1c41a94b62 | ||
|
|
2e66aafe20 | ||
|
|
55074bda76 | ||
|
|
de65bec2b7 | ||
|
|
7664dd0de3 | ||
|
|
019a88ced4 | ||
|
|
72de11abcc | ||
|
|
d71a4ebffc | ||
|
|
1089ab43bf |
@@ -0,0 +1,138 @@
|
||||
env:
|
||||
IMAGE_VERSION: "py3.12-latest"
|
||||
BUILDKITE_CLEAN_CHECKOUT: true
|
||||
|
||||
steps:
|
||||
- label: "pre-commit"
|
||||
command: ".buildkite/scripts/pre_commit.sh"
|
||||
agents:
|
||||
queue: "default"
|
||||
|
||||
- wait
|
||||
|
||||
- label: "Trigger Tests"
|
||||
plugins:
|
||||
- monorepo-diff#v1.4.0:
|
||||
diff: "git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD"
|
||||
watch:
|
||||
- path:
|
||||
- "fastvideo/v1/models/encoders/**"
|
||||
- "fastvideo/v1/models/loader/**"
|
||||
- "fastvideo/v1/tests/encoders/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Encoder Tests"
|
||||
env:
|
||||
- TEST_TYPE=encoder
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/models/vaes/**"
|
||||
- "fastvideo/v1/models/loader/**"
|
||||
- "fastvideo/v1/tests/vaes/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "VAE Tests"
|
||||
env:
|
||||
- TEST_TYPE=vae
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/models/dits/**"
|
||||
- "fastvideo/v1/models/loader/**"
|
||||
- "fastvideo/v1/tests/transformers/**"
|
||||
- "fastvideo/v1/layers/**"
|
||||
- "fastvideo/v1/attention/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Transformer Tests"
|
||||
env:
|
||||
- TEST_TYPE=transformer
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**/*.py"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "SSIM Tests"
|
||||
env:
|
||||
- TEST_TYPE=ssim
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Training Tests"
|
||||
env:
|
||||
- TEST_TYPE=training
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**"
|
||||
- "csrc/attn/vsa/**"
|
||||
- "csrc/attn/tk/**"
|
||||
- "csrc/attn/setup_vsa.py"
|
||||
- "csrc/attn/config_vsa.py"
|
||||
- "csrc/attn/vsa.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Training Tests VSA"
|
||||
env:
|
||||
- TEST_TYPE=training_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/**"
|
||||
- "csrc/attn/st_attn/**"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/attn/config_sta.py"
|
||||
- "csrc/attn/st_attn.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Inference Tests STA"
|
||||
env:
|
||||
- TEST_TYPE=inference_sta
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/st_attn/**"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
- "csrc/attn/config_sta.py"
|
||||
- "csrc/attn/st_attn.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests STA"
|
||||
env:
|
||||
- TEST_TYPE=precision_sta
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "csrc/attn/vsa/**"
|
||||
- "csrc/attn/tk/**"
|
||||
- "csrc/attn/setup_vsa.py"
|
||||
- "csrc/attn/config_vsa.py"
|
||||
- "csrc/attn/vsa.cpp"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
config:
|
||||
command: "timeout 15m .buildkite/scripts/pr_test.sh"
|
||||
label: "Precision Tests VSA"
|
||||
env:
|
||||
- TEST_TYPE=precision_vsa
|
||||
agents:
|
||||
queue: "default"
|
||||
Executable
+117
@@ -0,0 +1,117 @@
|
||||
#!/bin/bash
|
||||
set -uo pipefail
|
||||
|
||||
log() {
|
||||
echo "[$(date '+%Y-%m-%d %H:%M:%S')] $1"
|
||||
}
|
||||
|
||||
log "=== Starting Modal test execution ==="
|
||||
|
||||
# Change to the project directory
|
||||
cd "$(dirname "$0")/../.."
|
||||
PROJECT_ROOT=$(pwd)
|
||||
log "Project root: $PROJECT_ROOT"
|
||||
|
||||
# Install Modal if not available
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Modal not found, installing..."
|
||||
python3 -m pip install modal
|
||||
|
||||
# Verify installation
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Error: Failed to install modal. Please install it manually."
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
|
||||
log "modal version: $(python3 -m modal --version)"
|
||||
|
||||
# Set up Modal authentication using Buildkite secrets
|
||||
log "Setting up Modal authentication from Buildkite secrets..."
|
||||
MODAL_TOKEN_ID=$(buildkite-agent secret get modal_token_id)
|
||||
MODAL_TOKEN_SECRET=$(buildkite-agent secret get modal_token_secret)
|
||||
|
||||
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
|
||||
|
||||
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
|
||||
|
||||
if [ -n "$MODAL_TOKEN_ID" ] && [ -n "$MODAL_TOKEN_SECRET" ]; then
|
||||
log "Retrieved Modal credentials from Buildkite secrets"
|
||||
python3 -m modal token set --token-id "$MODAL_TOKEN_ID" --token-secret "$MODAL_TOKEN_SECRET" --profile buildkite-ci --activate --verify
|
||||
if [ $? -eq 0 ]; then
|
||||
log "Modal authentication successful"
|
||||
else
|
||||
log "Error: Failed to set Modal credentials"
|
||||
exit 1
|
||||
fi
|
||||
else
|
||||
log "Error: Could not retrieve Modal credentials from Buildkite secrets."
|
||||
log "Please ensure 'modal_token_id' and 'modal_token_secret' secrets are set in Buildkite."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
MODAL_TEST_FILE="fastvideo/v1/tests/modal/pr_test.py"
|
||||
|
||||
if [ -z "${TEST_TYPE:-}" ]; then
|
||||
log "Error: TEST_TYPE environment variable is not set"
|
||||
exit 1
|
||||
fi
|
||||
log "Test type: $TEST_TYPE"
|
||||
|
||||
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT BUILDKITE_PULL_REQUEST=$BUILDKITE_PULL_REQUEST IMAGE_VERSION=$IMAGE_VERSION"
|
||||
|
||||
case "$TEST_TYPE" in
|
||||
"encoder")
|
||||
log "Running encoder tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
|
||||
;;
|
||||
"vae")
|
||||
log "Running VAE tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
|
||||
;;
|
||||
"transformer")
|
||||
log "Running transformer tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
|
||||
;;
|
||||
"ssim")
|
||||
log "Running SSIM tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
;;
|
||||
"training")
|
||||
log "Running training tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests"
|
||||
;;
|
||||
"training_vsa")
|
||||
log "Running training VSA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests_VSA"
|
||||
;;
|
||||
"inference_sta")
|
||||
log "Running inference STA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_STA"
|
||||
;;
|
||||
"precision_sta")
|
||||
log "Running precision STA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_STA"
|
||||
;;
|
||||
"precision_vsa")
|
||||
log "Running precision VSA tests..."
|
||||
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_VSA"
|
||||
;;
|
||||
*)
|
||||
log "Error: Unknown test type: $TEST_TYPE"
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
|
||||
log "Executing: $MODAL_COMMAND"
|
||||
eval "$MODAL_COMMAND"
|
||||
TEST_EXIT_CODE=$?
|
||||
|
||||
if [ $TEST_EXIT_CODE -eq 0 ]; then
|
||||
log "Modal test completed successfully"
|
||||
else
|
||||
log "Error: Modal test failed with exit code: $TEST_EXIT_CODE"
|
||||
fi
|
||||
|
||||
log "=== Test execution completed with exit code: $TEST_EXIT_CODE ==="
|
||||
exit $TEST_EXIT_CODE
|
||||
@@ -0,0 +1,40 @@
|
||||
#!/bin/bash
|
||||
set -uo pipefail
|
||||
|
||||
log() {
|
||||
echo "[$(date '+%Y-%m-%d %H:%M:%S')] $1"
|
||||
}
|
||||
|
||||
log "=== Starting pre-commit checks ==="
|
||||
|
||||
cd "$(dirname "$0")/../.."
|
||||
PROJECT_ROOT=$(pwd)
|
||||
log "Project root: $PROJECT_ROOT"
|
||||
|
||||
if ! python3 -m pre_commit --version &> /dev/null; then
|
||||
log "pre-commit not found, installing..."
|
||||
python3 -m pip install --user pre-commit==4.0.1
|
||||
|
||||
if ! python3 -m pre_commit --version &> /dev/null; then
|
||||
log "Error: Failed to install pre-commit."
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
|
||||
log "Pre-commit version: $(python3 -m pre_commit --version)"
|
||||
|
||||
log "Installing/updating pre-commit hooks..."
|
||||
python3 -m pre_commit install --install-hooks
|
||||
|
||||
log "Running pre-commit checks on all files..."
|
||||
python3 -m pre_commit run --all-files
|
||||
PRE_COMMIT_EXIT_CODE=$?
|
||||
|
||||
if [ $PRE_COMMIT_EXIT_CODE -eq 0 ]; then
|
||||
log "Pre-commit checks completed successfully"
|
||||
else
|
||||
log "Error: Pre-commit checks failed with exit code: $PRE_COMMIT_EXIT_CODE"
|
||||
fi
|
||||
|
||||
log "=== Pre-commit checks completed with exit code: $PRE_COMMIT_EXIT_CODE ==="
|
||||
exit $PRE_COMMIT_EXIT_CODE
|
||||
@@ -160,8 +160,7 @@ def execute_command(pod_id):
|
||||
setup_steps = [
|
||||
"tar -xzf /tmp/repo.tar.gz --no-same-owner -C /workspace/",
|
||||
f"cd /workspace/{repo_name}",
|
||||
"source /opt/conda/etc/profile.d/conda.sh",
|
||||
"conda activate fastvideo-dev",
|
||||
"source $HOME/.local/bin/env && source /opt/venv/bin/activate",
|
||||
args.test_command
|
||||
]
|
||||
|
||||
|
||||
+161
-21
@@ -12,13 +12,11 @@ on:
|
||||
paths:
|
||||
- "fastvideo/**/*.py"
|
||||
- ".github/workflows/pr-test.yml"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
- "csrc/**"
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
custom_image:
|
||||
description: "Custom image from this repository (default: fastvideo-dev:latest)"
|
||||
required: false
|
||||
default: "fastvideo-dev:latest"
|
||||
type: string
|
||||
run_encoder_test:
|
||||
description: "Run encoder-test"
|
||||
required: false
|
||||
@@ -44,6 +42,26 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_training_test_VSA:
|
||||
description: "Run training-test-VSA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_inference_test_STA:
|
||||
description: "Run inference-test-STA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_precision_test_STA:
|
||||
description: "Run precision-test-STA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_precision_test_VSA:
|
||||
description: "Run precision-test-VSA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_nightly_test:
|
||||
description: "Run nightly-test"
|
||||
required: false
|
||||
@@ -53,6 +71,7 @@ on:
|
||||
env:
|
||||
PYTHONUNBUFFERED: "1"
|
||||
|
||||
|
||||
concurrency:
|
||||
group: pr-test-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
@@ -70,28 +89,71 @@ jobs:
|
||||
vae-test: ${{ steps.filter.outputs.vae-test }}
|
||||
transformer-test: ${{ steps.filter.outputs.transformer-test }}
|
||||
training-test: ${{ steps.filter.outputs.training-test }}
|
||||
training-test-VSA: ${{ steps.filter.outputs.training-test-VSA }}
|
||||
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
|
||||
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
|
||||
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
id: filter
|
||||
with:
|
||||
filters: |
|
||||
# Define reusable path patterns
|
||||
common-paths: &common-paths
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
sta-kernel-paths: &sta-kernel-paths
|
||||
- 'csrc/attn/st_attn/**'
|
||||
- 'csrc/attn/setup_sta.py'
|
||||
- 'csrc/attn/config_sta.py'
|
||||
- 'csrc/attn/st_attn.cpp'
|
||||
vsa-kernel-paths: &vsa-kernel-paths
|
||||
- 'csrc/attn/vsa/**'
|
||||
- 'csrc/attn/tk/**'
|
||||
- 'csrc/attn/setup_vsa.py'
|
||||
- 'csrc/attn/config_vsa.py'
|
||||
- 'csrc/attn/vsa.cpp'
|
||||
vsa-paths: &vsa-paths
|
||||
- 'fastvideo/v1/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
|
||||
# Actual tests
|
||||
encoder-test:
|
||||
- 'fastvideo/v1/models/encoders/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/encoders/**'
|
||||
- *common-paths
|
||||
vae-test:
|
||||
- 'fastvideo/v1/models/vaes/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/vaes/**'
|
||||
- *common-paths
|
||||
transformer-test:
|
||||
- 'fastvideo/v1/models/dits/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/transformers/**'
|
||||
- 'fastvideo/v1/layers/**'
|
||||
- 'fastvideo/v1/attention/**'
|
||||
- *common-paths
|
||||
training-test:
|
||||
- 'fastvideo/v1/**'
|
||||
- *common-paths
|
||||
training-test-VSA:
|
||||
- 'fastvideo/v1/**'
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
inference-test-STA:
|
||||
- 'fastvideo/v1/**'
|
||||
- *common-paths
|
||||
- *sta-kernel-paths
|
||||
precision-test-STA:
|
||||
- *common-paths
|
||||
- *sta-kernel-paths
|
||||
precision-test-VSA:
|
||||
- *common-paths
|
||||
- *vsa-kernel-paths
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
@@ -104,8 +166,8 @@ jobs:
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -122,8 +184,8 @@ jobs:
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -140,8 +202,8 @@ jobs:
|
||||
gpu_type: "NVIDIA L40S"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -150,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:
|
||||
@@ -168,7 +229,7 @@ jobs:
|
||||
volume_size: 200
|
||||
disk_size: 200
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:${{ matrix.python-version.tag }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
|
||||
timeout_minutes: 60
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -177,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:
|
||||
@@ -186,8 +247,86 @@ jobs:
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/training -srP"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/Vanilla -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
training-test-VSA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test-VSA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test_VSA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "training-test-VSA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 2
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/VSA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
inference-test-STA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.inference-test-STA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_inference_test_STA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "inference-test-STA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 2
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/inference/STA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
precision-test-STA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-STA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_STA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "precision-test-STA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_sta.py"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
precision-test-VSA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-VSA == 'true') ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_VSA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "precision-test-VSA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_block_sparse.py"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -203,8 +342,8 @@ jobs:
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -212,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:
|
||||
@@ -229,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
|
||||
|
||||
@@ -15,6 +15,11 @@ With FastVideo's optimizations, you can achieve more than 3x inference improveme
|
||||
<img src=assets/perf.png width="90%"/>
|
||||
</div>
|
||||
|
||||
## NEWS
|
||||
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
|
||||
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
|
||||
- ```2025/02/18```: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
|
||||
|
||||
## Key Features
|
||||
|
||||
FastVideo has the following features:
|
||||
@@ -91,7 +96,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
|
||||
|
||||
## Distillation and Finetuning
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/training/distillation.html)
|
||||
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetuning.html)
|
||||
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html)
|
||||
|
||||
## 📑 Development Plan
|
||||
|
||||
@@ -111,7 +116,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/developer_guide/overview.html)
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview.html)
|
||||
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects:
|
||||
@@ -128,6 +133,15 @@ We thank MBZUAI and [Anyscale](https://www.anyscale.com/) for their support thro
|
||||
If you use FastVideo for your research, please cite our paper:
|
||||
|
||||
```bibtex
|
||||
@misc{zhang2025vsafastervideodiffusion,
|
||||
title={VSA: Faster Video Diffusion with Trainable Sparse Attention},
|
||||
author={Peiyuan Zhang and Haofeng Huang and Yongqi Chen and Will Lin and Zhengzhong Liu and Ion Stoica and Eric Xing and Hao Zhang},
|
||||
year={2025},
|
||||
eprint={2505.13389},
|
||||
archivePrefix={arXiv},
|
||||
primaryClass={cs.CV},
|
||||
url={https://arxiv.org/abs/2505.13389},
|
||||
}
|
||||
@misc{zhang2025fastvideogenerationsliding,
|
||||
title={Fast Video Generation with Sliding Tile Attention},
|
||||
author={Peiyuan Zhang and Yongqi Chen and Runlong Su and Hangliang Ding and Ion Stoica and Zhenghong Liu and Hao Zhang},
|
||||
|
||||
+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();
|
||||
}
|
||||
@@ -1,7 +1,9 @@
|
||||
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
|
||||
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
@@ -9,17 +11,25 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
zsh \
|
||||
vim \
|
||||
curl \
|
||||
gcc-11 \
|
||||
g++-11 \
|
||||
clang-11 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
|
||||
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
|
||||
rm Miniconda3-latest-Linux-x86_64.sh
|
||||
# Set up C++20 compilers for ThunderKittens
|
||||
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
ENV PATH=/opt/conda/bin:$PATH
|
||||
# Set CUDA environment variables
|
||||
ENV CUDA_HOME=/usr/local/cuda-12.8
|
||||
ENV PATH=${CUDA_HOME}/bin:${PATH}
|
||||
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
RUN conda create --name fastvideo-dev python=3.12.9 -y
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
# Install uv and source its environment
|
||||
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
|
||||
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject.toml ./
|
||||
@@ -27,22 +37,36 @@ COPY pyproject.toml ./
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
|
||||
conda clean -afy
|
||||
# Create and activate virtual environment with specific Python version and seed
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.12 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir flash-attn==2.8.0.post2 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Remove authentication headers
|
||||
RUN git config --unset-all http.https://github.com/.extraheader || true
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_sta.py install
|
||||
|
||||
# Set up automatic conda environment activation for all shells
|
||||
RUN echo 'source /opt/conda/etc/profile.d/conda.sh' >> /root/.bashrc && \
|
||||
echo 'conda activate fastvideo-dev' >> /root/.bashrc && \
|
||||
# Ensure .bashrc is sourced for SSH login shells
|
||||
echo 'if [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_vsa.py install
|
||||
|
||||
EXPOSE 22
|
||||
@@ -1,7 +1,7 @@
|
||||
(sta-demo)=
|
||||
|
||||
# 🔍 Demo
|
||||
There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
This is is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
<div style="text-align: center;">
|
||||
<video controls width="800">
|
||||
@@ -9,3 +9,9 @@ There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
Your browser does not support the video tag.
|
||||
</video>
|
||||
</div>
|
||||
|
||||
You can run STA using the following command:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_STA.sh
|
||||
```
|
||||
|
||||
@@ -16,6 +16,13 @@ bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
## ⚡ Finetune with VSA
|
||||
Follow [data_preprocess.md](#v0-data-preprocess) to get parquet files for preproccessed latent, and then run:
|
||||
|
||||
```bash
|
||||
bash scripts/finetune/finetune_v1_VSA.sh
|
||||
```
|
||||
|
||||
## ⚡ Lora Finetune
|
||||
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
|
||||
@@ -5,7 +5,7 @@ export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
|
||||
export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
|
||||
|
||||
base_port=29503
|
||||
num_gpu=$(nvidia-smi --query-gpu=gpu_name --format=csv,noheader | wc -l)
|
||||
num_gpu=1
|
||||
gpu_ids=$(seq 0 $((num_gpu-1)))
|
||||
skip_time_steps=12
|
||||
|
||||
@@ -14,7 +14,7 @@ STA_mode="STA_searching"
|
||||
for i in $gpu_ids; do
|
||||
port=$((base_port+i))
|
||||
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
|
||||
--prompt_path ./assets/prompt_extend_${i}.txt \
|
||||
--prompt_path ./assets/prompt_${i}.txt \
|
||||
--output_path $output_path \
|
||||
--STA_mode $STA_mode &
|
||||
sleep 1
|
||||
@@ -27,7 +27,7 @@ STA_mode="STA_tuning"
|
||||
for i in $gpu_ids; do
|
||||
port=$((base_port+i))
|
||||
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
|
||||
--prompt_path ./assets/prompt_extend_${i}.txt \
|
||||
--prompt_path ./assets/prompt_${i}.txt \
|
||||
--output_path $output_path \
|
||||
--STA_mode $STA_mode \
|
||||
--skip_time_steps $skip_time_steps &
|
||||
|
||||
@@ -0,0 +1,10 @@
|
||||
This directory contain e2e examples scripts for finetuning Wan2.1 I2V.
|
||||
|
||||
Execute the following commands from `FastVideo/` to run training:
|
||||
|
||||
- Download crush-smol dataset:
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/download_dataset.sh`
|
||||
- Preprocess the videos and captions into latents:
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/preprocess_wan_data_i2v.sh`
|
||||
- Edit the following file and run finetuning:
|
||||
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/finetune_i2v.sh`
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,91 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_i2v/combined_parquet_dataset/"
|
||||
VALIDATION_DIR="data/crush-smol_processed_i2v/validation_parquet_dataset/"
|
||||
NUM_GPUS=8
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_i2v_finetune"
|
||||
--output_dir "$DATA_DIR/outputs/wan_i2v_finetune"
|
||||
--max_train_steps 2000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 8
|
||||
--tp_size 8
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_preprocessed_path "$VALIDATION_DIR"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,130 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=i2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --qos=hao
|
||||
#SBATCH --nodes=4
|
||||
#SBATCH --ntasks=4
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=i2v_output/i2v_%j.out
|
||||
#SBATCH --error=i2v_output/i2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_i2v/combined_parquet_dataset/"
|
||||
VALIDATION_DIR="data/crush-smol_processed_i2v/validation_parquet_dataset/"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_i2v_finetune
|
||||
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune_2n"
|
||||
--max_train_steps=2000
|
||||
--train_batch_size=2
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps=1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size $NUM_GPUS
|
||||
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 10
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_preprocessed_path "$VALIDATION_DIR"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate=1e-5
|
||||
--mixed_precision="bf16"
|
||||
--checkpointing_steps=1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,26 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_i2v/"
|
||||
VALIDATION_PATH="examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation.json"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 77 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--model_type $MODEL_TYPE \
|
||||
--train_fps 16 \
|
||||
--validation_dataset_file $VALIDATION_PATH \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "i2v"
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -0,0 +1,10 @@
|
||||
This directory contain e2e examples scripts for finetuning Wan2.1 T2v.
|
||||
|
||||
Execute the following commands from `FastVideo/` to run training:
|
||||
|
||||
- Download crush-smol dataset:
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/download_dataset.sh`
|
||||
- Preprocess the videos and captions into latents:
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/preprocess_wan_data_t2v.sh`
|
||||
- Edit the following file and run finetuning:
|
||||
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/finetune_t2v.sh`
|
||||
@@ -0,0 +1,3 @@
|
||||
#!/bin/bash
|
||||
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,90 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
|
||||
VALIDATION_DIR="data/crush-smol_processed_t2v/validation_parquet_dataset/"
|
||||
NUM_GPUS=4
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_t2v_finetune"
|
||||
--output_dir "outputs/wan_t2v_finetune"
|
||||
--max_train_steps 5000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 8
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size $NUM_GPUS
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path $DATA_DIR
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_preprocessed_path $VALIDATION_DIR
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 5e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 6000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,127 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=t2v
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --qos=hao
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=t2v_output/t2v_%j.out
|
||||
#SBATCH --error=t2v_output/t2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
echo "MASTER_ADDR: $MASTER_ADDR"
|
||||
echo "NODE_RANK: $NODE_RANK"
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
|
||||
VALIDATION_DIR="data/crush-smol_processed_t2v/validation_parquet_dataset/"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name wan_t2v_finetune
|
||||
--output_dir="outputs/wan_t2v_finetune"
|
||||
--max_train_steps=1000
|
||||
--train_batch_size=4
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps=1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 4
|
||||
--tp_size 4
|
||||
--hsdp_replicate_dim 2
|
||||
--hsdp_shard_dim 4
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 10
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_preprocessed_path "$VALIDATION_DIR"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate=5e-5
|
||||
--mixed_precision="bf16"
|
||||
--checkpointing_steps=500
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.1
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
)
|
||||
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/v1/training/wan_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,26 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v/"
|
||||
VALIDATION_PATH="examples/training/finetune/wan_t2v_1_3b/crush_smol/validation.json"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 77 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--model_type $MODEL_TYPE \
|
||||
--train_fps 16 \
|
||||
--validation_dataset_file $VALIDATION_PATH \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--video_length_tolerance_range 5 \
|
||||
--preprocess_task "t2v"
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -6,6 +6,8 @@ from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.v1.utils import dict_to_3d_list
|
||||
|
||||
|
||||
def configure_sta(mode: str = 'STA_searching',
|
||||
layer_num: int = 40,
|
||||
@@ -349,21 +351,6 @@ def select_best_mask_strategy(
|
||||
return best_mask_strategy, overall_sparsity, strategy_counts
|
||||
|
||||
|
||||
def dict_to_3d_list(mask_strategy: Optional[Dict[str, List[int]]],
|
||||
t_max: int = 50,
|
||||
l_max: int = 60,
|
||||
h_max: int = 24) -> List[List[List[Optional[List[int]]]]]:
|
||||
result: List[List[List[Optional[List[int]]]]] = [[[
|
||||
None for _ in range(h_max)
|
||||
] for _ in range(l_max)] for _ in range(t_max)]
|
||||
if mask_strategy is None:
|
||||
return result
|
||||
for key, value in mask_strategy.items():
|
||||
t, layer_idx, h = map(int, key.split('_'))
|
||||
result[t][layer_idx][h] = value
|
||||
return result
|
||||
|
||||
|
||||
def save_mask_search_results(
|
||||
mask_search_final_result: List[Dict[str, List[float]]],
|
||||
prompt: str,
|
||||
|
||||
@@ -10,8 +10,8 @@ from fastvideo.v1.attention.selector import get_attn_backend
|
||||
|
||||
__all__ = [
|
||||
"DistributedAttention",
|
||||
"DistributedAttention_VSA",
|
||||
"LocalAttention",
|
||||
"DistributedAttention_VSA",
|
||||
"AttentionBackend",
|
||||
"AttentionMetadata",
|
||||
"AttentionMetadataBuilder",
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Type
|
||||
from typing import Any, List, Optional, Type
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
@@ -17,33 +17,11 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.utils import dict_to_3d_list
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
# TODO(will-refactor): move this to a utils file
|
||||
def dict_to_3d_list(
|
||||
mask_strategy: Dict[str,
|
||||
Any]) -> List[List[List[Optional[torch.Tensor]]]]:
|
||||
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
|
||||
|
||||
max_timesteps_idx = max(
|
||||
timesteps_idx for timesteps_idx, layer_idx, head_idx in indices) + 1
|
||||
max_layer_idx = max(layer_idx
|
||||
for timesteps_idx, layer_idx, head_idx in indices) + 1
|
||||
max_head_idx = max(head_idx
|
||||
for timesteps_idx, layer_idx, head_idx in indices) + 1
|
||||
|
||||
result = [[[None for _ in range(max_head_idx)]
|
||||
for _ in range(max_layer_idx)] for _ in range(max_timesteps_idx)]
|
||||
|
||||
for key, value in mask_strategy.items():
|
||||
timesteps_idx, layer_idx, head_idx = map(int, key.split('_'))
|
||||
result[timesteps_idx][layer_idx][head_idx] = value
|
||||
|
||||
return result
|
||||
|
||||
|
||||
class RangeDict(dict):
|
||||
|
||||
def __getitem__(self, item: int) -> str:
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
from dataclasses import asdict, dataclass, field, fields
|
||||
from enum import Enum
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
import torch
|
||||
@@ -16,6 +17,15 @@ from fastvideo.v1.utils import (FlexibleArgumentParser, StoreBoolean,
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class STA_Mode(str, Enum):
|
||||
"""STA (Sliding Tile Attention) modes."""
|
||||
STA_INFERENCE = "STA_inference"
|
||||
STA_SEARCHING = "STA_searching"
|
||||
STA_TUNING = "STA_tuning"
|
||||
STA_TUNING_CFG = "STA_tuning_cfg"
|
||||
NONE = None
|
||||
|
||||
|
||||
def preprocess_text(prompt: str) -> str:
|
||||
return prompt
|
||||
|
||||
@@ -42,7 +52,7 @@ class PipelineConfig:
|
||||
|
||||
# VAE configuration
|
||||
vae_config: VAEConfig = field(default_factory=VAEConfig)
|
||||
vae_precision: str = "fp16"
|
||||
vae_precision: str = "fp32"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = True
|
||||
|
||||
@@ -76,7 +86,7 @@ class PipelineConfig:
|
||||
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
STA_mode: Optional[str] = None
|
||||
STA_mode: STA_Mode = STA_Mode.STA_INFERENCE
|
||||
skip_time_steps: int = 15
|
||||
|
||||
# Compilation
|
||||
|
||||
@@ -50,7 +50,7 @@ class WanT2V480PConfig(PipelineConfig):
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
vae_precision: str = "fp32"
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp32", ))
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
"use_cpu_offload": true,
|
||||
"disable_autocast": false,
|
||||
"precision": "bf16",
|
||||
"vae_precision": "fp16",
|
||||
"vae_precision": "fp32",
|
||||
"vae_tiling": false,
|
||||
"vae_sp": false,
|
||||
"vae_config": {
|
||||
|
||||
@@ -4,7 +4,7 @@
|
||||
"use_cpu_offload": true,
|
||||
"disable_autocast": false,
|
||||
"precision": "bf16",
|
||||
"vae_precision": "fp16",
|
||||
"vae_precision": "fp32",
|
||||
"vae_tiling": false,
|
||||
"vae_sp": false,
|
||||
"vae_config": {
|
||||
|
||||
@@ -1,19 +1,17 @@
|
||||
import os
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from torchvision import transforms
|
||||
from torchvision.transforms import Lambda
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.v1.dataset.t2v_datasets import T2V_dataset
|
||||
from fastvideo.v1.dataset.parquet_dataset_map_style import (
|
||||
build_parquet_map_style_dataloader)
|
||||
from fastvideo.v1.dataset.preprocessing_datasets import (
|
||||
VideoCaptionMergedDataset)
|
||||
from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
|
||||
from .parquet_dataset_map_style import build_parquet_map_style_dataloader
|
||||
|
||||
__all__ = ["build_parquet_map_style_dataloader"]
|
||||
from fastvideo.v1.dataset.validation_dataset import ValidationDataset
|
||||
|
||||
|
||||
def getdataset(args, start_idx=0) -> T2V_dataset:
|
||||
def getdataset(args) -> VideoCaptionMergedDataset:
|
||||
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
|
||||
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
|
||||
resize_topcrop = [
|
||||
@@ -31,15 +29,14 @@ def getdataset(args, start_idx=0) -> T2V_dataset:
|
||||
*resize_topcrop,
|
||||
norm_fun,
|
||||
])
|
||||
tokenizer_path = os.path.join(args.model_path, "tokenizer")
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
cache_dir=args.cache_dir)
|
||||
if args.dataset == "t2v":
|
||||
return T2V_dataset(args,
|
||||
transform=transform,
|
||||
temporal_sample=temporal_sample,
|
||||
tokenizer=tokenizer,
|
||||
transform_topcrop=transform_topcrop,
|
||||
start_idx=start_idx)
|
||||
return VideoCaptionMergedDataset(data_merge_path=args.data_merge_path,
|
||||
args=args,
|
||||
transform=transform,
|
||||
temporal_sample=temporal_sample,
|
||||
transform_topcrop=transform_topcrop)
|
||||
|
||||
raise NotImplementedError(args.dataset)
|
||||
|
||||
__all__ = [
|
||||
"build_parquet_map_style_dataloader", "ValidationDataset",
|
||||
"VideoCaptionMergedDataset"
|
||||
]
|
||||
|
||||
@@ -8,6 +8,7 @@ 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
|
||||
@@ -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,8 +158,9 @@ 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
|
||||
|
||||
|
||||
@@ -26,15 +26,47 @@ pyarrow_schema_i2v = pa.schema([
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
pa.field("text_attention_mask_bytes", pa.binary()),
|
||||
# e.g., [SeqLen]
|
||||
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bool' or 'int8'
|
||||
pa.field("text_attention_mask_dtype", pa.string()),
|
||||
#I2V
|
||||
pa.field("clip_feature_bytes", pa.binary()),
|
||||
pa.field("clip_feature_shape", pa.list_(pa.int64())),
|
||||
pa.field("clip_feature_dtype", pa.string()),
|
||||
pa.field("first_frame_latent_bytes", pa.binary()),
|
||||
pa.field("first_frame_latent_shape", pa.list_(pa.int64())),
|
||||
pa.field("first_frame_latent_dtype", pa.string()),
|
||||
# I2V Validation
|
||||
pa.field("pil_image_bytes", pa.binary()),
|
||||
pa.field("pil_image_shape", pa.list_(pa.int64())),
|
||||
pa.field("pil_image_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
pa.field("media_type", pa.string()), # 'image' or 'video'
|
||||
pa.field("width", pa.int64()),
|
||||
pa.field("height", pa.int64()),
|
||||
# -- Video-specific (can be null/default for images) ---
|
||||
# Number of frames processed (e.g., 1 for image, N for video)
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
pyarrow_schema_i2v_validation = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
#I2V
|
||||
pa.field("clip_feature_bytes", pa.binary()),
|
||||
pa.field("clip_feature_shape", pa.list_(pa.int64())),
|
||||
pa.field("clip_feature_dtype", pa.string()),
|
||||
# I2V Validation
|
||||
pa.field("pil_image_bytes", pa.binary()),
|
||||
pa.field("pil_image_shape", pa.list_(pa.int64())),
|
||||
pa.field("pil_image_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
@@ -64,11 +96,6 @@ pyarrow_schema_t2v = pa.schema([
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
pa.field("text_attention_mask_bytes", pa.binary()),
|
||||
# e.g., [SeqLen]
|
||||
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bool' or 'int8'
|
||||
pa.field("text_attention_mask_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
@@ -80,4 +107,26 @@ pyarrow_schema_t2v = pa.schema([
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
])
|
||||
|
||||
pyarrow_schema_t2v_validation = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
pa.field("media_type", pa.string()), # 'image' or 'video'
|
||||
pa.field("width", pa.int64()),
|
||||
pa.field("height", pa.int64()),
|
||||
# -- Video-specific (can be null/default for images) ---
|
||||
# Number of frames processed (e.g., 1 for image, N for video)
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
@@ -4,6 +4,7 @@ import random
|
||||
from typing import Dict, List, Tuple
|
||||
|
||||
import numpy as np
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
import tqdm
|
||||
@@ -70,10 +71,12 @@ class LatentsParquetIterStyleDataset(IterableDataset):
|
||||
drop_last: bool = True,
|
||||
text_padding_length: int = 512,
|
||||
seed: int = 42,
|
||||
read_batch_size: int = 32):
|
||||
read_batch_size: int = 32,
|
||||
parquet_schema: pa.Schema = None):
|
||||
super().__init__()
|
||||
self.path = str(path)
|
||||
self.batch_size = batch_size
|
||||
self.parquet_schema = parquet_schema
|
||||
self.cfg_rate = cfg_rate
|
||||
self.text_padding_length = text_padding_length
|
||||
self.seed = seed
|
||||
|
||||
@@ -3,6 +3,7 @@ import os
|
||||
import pickle
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
# Torch in general
|
||||
import torch
|
||||
@@ -11,7 +12,7 @@ import tqdm
|
||||
from torch.utils.data import Dataset, Sampler
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.v1.dataset.utils import collate_latents_embs_masks
|
||||
from fastvideo.v1.dataset.utils import collate_rows_from_parquet_schema
|
||||
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
|
||||
get_world_size)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -184,13 +185,12 @@ 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"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path: str,
|
||||
batch_size: int,
|
||||
parquet_schema: pa.Schema,
|
||||
cfg_rate: float = 0.0,
|
||||
seed: int = 42,
|
||||
drop_last: bool = True,
|
||||
@@ -200,24 +200,12 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
super().__init__()
|
||||
self.path = path
|
||||
self.cfg_rate = cfg_rate
|
||||
if cfg_rate > 0.0:
|
||||
raise ValueError(
|
||||
"cfg_rate > 0.0 is not supported for now because it will trigger bug when num_data_workers > 0"
|
||||
)
|
||||
self.parquet_schema = parquet_schema
|
||||
logger.info("Initializing LatentsParquetMapStyleDataset with path: %s",
|
||||
path)
|
||||
self.parquet_files, self.lengths = get_parquet_files_and_length(path)
|
||||
self.batch = batch_size
|
||||
self.text_padding_length = text_padding_length
|
||||
self._cols = [
|
||||
"vae_latent_bytes",
|
||||
"vae_latent_shape",
|
||||
"text_embedding_bytes",
|
||||
"text_embedding_shape",
|
||||
"text_embedding_dtype",
|
||||
"height",
|
||||
"width",
|
||||
]
|
||||
self.sampler = DP_SP_BatchSampler(
|
||||
batch_size=batch_size,
|
||||
dataset_size=sum(self.lengths),
|
||||
@@ -232,7 +220,7 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
len(self.parquet_files), sum(self.lengths))
|
||||
|
||||
def get_validation_negative_prompt(
|
||||
self) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, str]:
|
||||
self) -> tuple[torch.Tensor, torch.Tensor, str]:
|
||||
"""
|
||||
Get the negative prompt for validation.
|
||||
This method ensures the negative prompt is loaded and cached properly.
|
||||
@@ -246,19 +234,23 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
row_dict = read_row_from_parquet_file([file_path], row_idx,
|
||||
[self.lengths[0]])
|
||||
|
||||
all_latents_list, all_embs_list, all_masks_list, caption_text_list = collate_latents_embs_masks(
|
||||
[row_dict], self.text_padding_length, self.keys)
|
||||
all_latents, all_embs, all_masks, caption_text = all_latents_list[
|
||||
0], all_embs_list[0], all_masks_list[0], caption_text_list[0]
|
||||
# add batch dimension
|
||||
if len(all_embs.shape) == 2:
|
||||
all_embs = all_embs.unsqueeze(0)
|
||||
if len(all_masks.shape) == 1:
|
||||
all_masks = all_masks.unsqueeze(0).unsqueeze(0)
|
||||
return all_latents, all_embs, all_masks, caption_text
|
||||
batch = collate_rows_from_parquet_schema([row_dict],
|
||||
self.parquet_schema,
|
||||
self.text_padding_length,
|
||||
cfg_rate=0.0)
|
||||
negative_prompt = batch['info_list'][0]['prompt']
|
||||
negative_prompt_embedding = batch['text_embedding']
|
||||
negative_prompt_attention_mask = batch['text_attention_mask']
|
||||
if len(negative_prompt_embedding.shape) == 2:
|
||||
negative_prompt_embedding = negative_prompt_embedding.unsqueeze(0)
|
||||
if len(negative_prompt_attention_mask.shape) == 1:
|
||||
negative_prompt_attention_mask = negative_prompt_attention_mask.unsqueeze(
|
||||
0).unsqueeze(0)
|
||||
|
||||
return negative_prompt_embedding, negative_prompt_attention_mask, negative_prompt
|
||||
|
||||
# PyTorch calls this ONLY because the batch_sampler yields a list
|
||||
def __getitems__(self, indices: List[int]):
|
||||
def __getitems__(self, indices: List[int]) -> Dict[str, Any]:
|
||||
"""
|
||||
Batch fetch using read_row_from_parquet_file for each index.
|
||||
"""
|
||||
@@ -267,9 +259,11 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
for idx in indices
|
||||
]
|
||||
|
||||
all_latents, all_embs, all_masks, caption_text = collate_latents_embs_masks(
|
||||
rows, self.text_padding_length, self.keys)
|
||||
return all_latents, all_embs, all_masks, caption_text
|
||||
batch = collate_rows_from_parquet_schema(rows,
|
||||
self.parquet_schema,
|
||||
self.text_padding_length,
|
||||
cfg_rate=self.cfg_rate)
|
||||
return batch
|
||||
|
||||
def __len__(self):
|
||||
return sum(self.lengths)
|
||||
@@ -286,6 +280,7 @@ def build_parquet_map_style_dataloader(
|
||||
path,
|
||||
batch_size,
|
||||
num_data_workers,
|
||||
parquet_schema,
|
||||
cfg_rate=0.0,
|
||||
drop_last=True,
|
||||
drop_first_row=False,
|
||||
@@ -298,6 +293,7 @@ def build_parquet_map_style_dataloader(
|
||||
drop_last=drop_last,
|
||||
drop_first_row=drop_first_row,
|
||||
text_padding_length=text_padding_length,
|
||||
parquet_schema=parquet_schema,
|
||||
seed=seed)
|
||||
|
||||
loader = StatefulDataLoader(
|
||||
|
||||
@@ -0,0 +1,615 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import Counter
|
||||
from dataclasses import dataclass
|
||||
from os.path import join as opj
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class PreprocessBatch:
|
||||
"""
|
||||
Batch information for dataset processing stages.
|
||||
|
||||
This class holds all the information about a video-caption or image-caption pair
|
||||
as it moves through the processing pipeline. Fields are populated by different stages.
|
||||
"""
|
||||
# Raw metadata
|
||||
path: str
|
||||
cap: Union[str, List[str]]
|
||||
resolution: Optional[Dict] = None
|
||||
fps: Optional[float] = None
|
||||
duration: Optional[float] = None
|
||||
|
||||
# Processed metadata
|
||||
num_frames: Optional[int] = None
|
||||
sample_frame_index: Optional[List[int]] = None
|
||||
sample_num_frames: Optional[int] = None
|
||||
|
||||
# Processed data
|
||||
pixel_values: Optional[torch.Tensor] = None
|
||||
text: Optional[str] = None
|
||||
input_ids: Optional[torch.Tensor] = None
|
||||
cond_mask: Optional[torch.Tensor] = None
|
||||
|
||||
@property
|
||||
def is_video(self) -> bool:
|
||||
"""Check if this is a video item."""
|
||||
return self.path.endswith(".mp4")
|
||||
|
||||
@property
|
||||
def is_image(self) -> bool:
|
||||
"""Check if this is an image item."""
|
||||
return self.path.endswith(".jpg")
|
||||
|
||||
|
||||
class DatasetStage(ABC):
|
||||
"""
|
||||
Abstract base class for dataset processing stages.
|
||||
|
||||
Similar to PipelineStage but designed for dataset preprocessing operations.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Process the dataset batch.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch to process
|
||||
**kwargs: Additional processing parameters
|
||||
|
||||
Returns:
|
||||
Processed batch
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class DatasetFilterStage(ABC):
|
||||
"""
|
||||
Abstract base class for dataset filtering stages.
|
||||
|
||||
These stages can filter out items during metadata processing.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
|
||||
"""
|
||||
Check if batch should be kept.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch to check
|
||||
**kwargs: Additional parameters
|
||||
|
||||
Returns:
|
||||
True if batch should be kept, False otherwise
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Process the dataset batch (for non-filtering operations).
|
||||
|
||||
Args:
|
||||
batch: Dataset batch to process
|
||||
**kwargs: Additional processing parameters
|
||||
|
||||
Returns:
|
||||
Processed batch
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class DataValidationStage(DatasetFilterStage):
|
||||
"""Stage for validating data items."""
|
||||
|
||||
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
|
||||
"""
|
||||
Validate data item.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch to validate
|
||||
|
||||
Returns:
|
||||
True if valid, False if invalid
|
||||
"""
|
||||
# Check for caption
|
||||
if batch.cap is None:
|
||||
return False
|
||||
|
||||
if batch.is_video:
|
||||
# Validate video-specific fields
|
||||
if batch.duration is None or batch.fps is None:
|
||||
return False
|
||||
elif not batch.is_image:
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""Process does nothing for validation - filtering is handled by should_keep."""
|
||||
return batch
|
||||
|
||||
|
||||
class ResolutionFilterStage(DatasetFilterStage):
|
||||
"""Stage for filtering data items based on resolution constraints."""
|
||||
|
||||
def __init__(self,
|
||||
max_h_div_w_ratio: float = 17 / 16,
|
||||
min_h_div_w_ratio: float = 8 / 16,
|
||||
max_height: int = 1024,
|
||||
max_width: int = 1024):
|
||||
self.max_h_div_w_ratio = max_h_div_w_ratio
|
||||
self.min_h_div_w_ratio = min_h_div_w_ratio
|
||||
self.max_height = max_height
|
||||
self.max_width = max_width
|
||||
|
||||
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
|
||||
"""
|
||||
Check if data item passes resolution filtering.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch with resolution information
|
||||
|
||||
Returns:
|
||||
True if passes filter, False otherwise
|
||||
"""
|
||||
# Only apply to videos
|
||||
if not batch.is_video:
|
||||
return True
|
||||
|
||||
if batch.resolution is None:
|
||||
return False
|
||||
|
||||
height = batch.resolution.get("height", None)
|
||||
width = batch.resolution.get("width", None)
|
||||
if height is None or width is None:
|
||||
return False
|
||||
|
||||
# Check aspect ratio
|
||||
aspect = self.max_height / self.max_width
|
||||
hw_aspect_thr = 1.5
|
||||
|
||||
return self.filter_resolution(
|
||||
height,
|
||||
width,
|
||||
max_h_div_w_ratio=hw_aspect_thr * aspect,
|
||||
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
|
||||
)
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""Process does nothing for resolution filtering - filtering is handled by should_keep."""
|
||||
return batch
|
||||
|
||||
def filter_resolution(self, h: int, w: int, max_h_div_w_ratio: float,
|
||||
min_h_div_w_ratio: float) -> bool:
|
||||
"""Filter based on height/width ratio."""
|
||||
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
|
||||
|
||||
|
||||
class FrameSamplingStage(DatasetFilterStage):
|
||||
"""Stage for temporal frame sampling and indexing."""
|
||||
|
||||
def __init__(self,
|
||||
num_frames: int,
|
||||
train_fps: int,
|
||||
speed_factor: int = 1,
|
||||
video_length_tolerance_range: float = 5.0,
|
||||
drop_short_ratio: float = 0.0):
|
||||
self.num_frames = num_frames
|
||||
self.train_fps = train_fps
|
||||
self.speed_factor = speed_factor
|
||||
self.video_length_tolerance_range = video_length_tolerance_range
|
||||
self.drop_short_ratio = drop_short_ratio
|
||||
|
||||
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
|
||||
"""
|
||||
Check if video should be kept based on length constraints.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch
|
||||
|
||||
Returns:
|
||||
True if should be kept, False otherwise
|
||||
"""
|
||||
if batch.is_image:
|
||||
return True
|
||||
|
||||
if batch.duration is None or batch.fps is None:
|
||||
return False
|
||||
|
||||
num_frames = math.ceil(batch.fps * batch.duration)
|
||||
|
||||
# Check if video is too long
|
||||
if (num_frames / batch.fps > self.video_length_tolerance_range *
|
||||
(self.num_frames / self.train_fps * self.speed_factor)):
|
||||
return False
|
||||
|
||||
# Resample frame indices to check length
|
||||
frame_interval = batch.fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, num_frames,
|
||||
frame_interval).astype(int)
|
||||
|
||||
# Filter short videos
|
||||
return not (len(frame_indices) < self.num_frames
|
||||
and random.random() < self.drop_short_ratio)
|
||||
|
||||
def process(self,
|
||||
batch: PreprocessBatch,
|
||||
temporal_sample_fn=None,
|
||||
**kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Process frame sampling for video data items.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch
|
||||
temporal_sample_fn: Function for temporal sampling
|
||||
|
||||
Returns:
|
||||
Updated batch with frame sampling info
|
||||
"""
|
||||
if batch.is_image:
|
||||
# For images, just add sample info
|
||||
batch.sample_frame_index = [0]
|
||||
batch.sample_num_frames = 1
|
||||
return batch
|
||||
|
||||
assert batch.duration is not None and batch.fps is not None
|
||||
batch.num_frames = math.ceil(batch.fps * batch.duration)
|
||||
|
||||
# Resample frame indices
|
||||
frame_interval = batch.fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, batch.num_frames,
|
||||
frame_interval).astype(int)
|
||||
|
||||
# Temporal crop if too long
|
||||
if len(frame_indices
|
||||
) > self.num_frames and temporal_sample_fn is not None:
|
||||
begin_index, end_index = temporal_sample_fn(len(frame_indices))
|
||||
frame_indices = frame_indices[begin_index:end_index]
|
||||
|
||||
batch.sample_frame_index = frame_indices.tolist()
|
||||
batch.sample_num_frames = len(frame_indices)
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
class VideoTransformStage(DatasetStage):
|
||||
"""Stage for video data transformation."""
|
||||
|
||||
def __init__(self, transform) -> None:
|
||||
self.transform = transform
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Transform video data.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch with video information
|
||||
|
||||
Returns:
|
||||
Batch with transformed video tensor
|
||||
"""
|
||||
if not batch.is_video:
|
||||
return batch
|
||||
|
||||
assert os.path.exists(batch.path), f"file {batch.path} do not exist!"
|
||||
assert batch.sample_frame_index is not None, "Frame indices must be set before transformation"
|
||||
|
||||
torchvision_video, _, metadata = torchvision.io.read_video(
|
||||
batch.path, output_format="TCHW")
|
||||
video = torchvision_video[batch.sample_frame_index]
|
||||
if self.transform is not None:
|
||||
video = self.transform(video)
|
||||
video = rearrange(video, "t c h w -> c t h w")
|
||||
video = video.to(torch.uint8)
|
||||
|
||||
h, w = video.shape[-2:]
|
||||
assert (
|
||||
h / w <= 17 / 16 and h / w >= 8 / 16
|
||||
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({batch.path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
|
||||
|
||||
video = video.float() / 127.5 - 1.0
|
||||
batch.pixel_values = video
|
||||
return batch
|
||||
|
||||
|
||||
class ImageTransformStage(DatasetStage):
|
||||
"""Stage for image data transformation."""
|
||||
|
||||
def __init__(self, transform, transform_topcrop) -> None:
|
||||
self.transform = transform
|
||||
self.transform_topcrop = transform_topcrop
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Transform image data.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch with image information
|
||||
|
||||
Returns:
|
||||
Batch with transformed image tensor
|
||||
"""
|
||||
if not batch.is_image:
|
||||
return batch
|
||||
|
||||
image = Image.open(batch.path).convert("RGB")
|
||||
image = torch.from_numpy(np.array(image))
|
||||
image = rearrange(image, "h w c -> c h w").unsqueeze(0)
|
||||
|
||||
if self.transform_topcrop is not None:
|
||||
image = self.transform_topcrop(image)
|
||||
elif self.transform is not None:
|
||||
image = self.transform(image)
|
||||
|
||||
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
|
||||
image = image.float() / 127.5 - 1.0
|
||||
batch.pixel_values = image
|
||||
return batch
|
||||
|
||||
|
||||
class TextEncodingStage(DatasetStage):
|
||||
"""Stage for text tokenization and encoding."""
|
||||
|
||||
def __init__(self, tokenizer, text_max_length: int, cfg_rate: float = 0.0):
|
||||
self.tokenizer = tokenizer
|
||||
self.text_max_length = text_max_length
|
||||
self.cfg_rate = cfg_rate
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Process text data.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch with caption information
|
||||
|
||||
Returns:
|
||||
Batch with encoded text information
|
||||
"""
|
||||
text = batch.cap
|
||||
if not isinstance(text, list):
|
||||
text = [text]
|
||||
text = [random.choice(text)]
|
||||
|
||||
text = text[0] if random.random() > self.cfg_rate else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
batch.text = text
|
||||
batch.input_ids = text_tokens_and_mask["input_ids"]
|
||||
batch.cond_mask = text_tokens_and_mask["attention_mask"]
|
||||
return batch
|
||||
|
||||
|
||||
class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
torch.distributed.checkpoint.stateful.Stateful):
|
||||
"""
|
||||
Merged dataset for video and caption data with stage-based processing.
|
||||
Assumes that data_merge_path is a txt file with the following format:
|
||||
<folder_path>,<json_file_path>
|
||||
|
||||
The folder should contain videos.
|
||||
|
||||
The json file should be a list of dictionaries with the following format:
|
||||
[
|
||||
{
|
||||
"path": "1gGQy4nxyUo-Scene-016.mp4",
|
||||
"resolution": {
|
||||
"width": 1920,
|
||||
"height": 1080
|
||||
},
|
||||
"size": 2439112,
|
||||
"fps": 25.0,
|
||||
"duration": 6.88,
|
||||
"num_frames": 172,
|
||||
"cap": [
|
||||
"A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open."
|
||||
]
|
||||
},
|
||||
...
|
||||
]
|
||||
|
||||
This dataset processes video and image data through a series of stages:
|
||||
- Data validation
|
||||
- Resolution filtering
|
||||
- Frame sampling
|
||||
- Transformation
|
||||
- Text encoding
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
data_merge_path: str,
|
||||
args,
|
||||
transform,
|
||||
temporal_sample,
|
||||
transform_topcrop,
|
||||
start_idx: int = 0):
|
||||
self.data_merge_path = data_merge_path
|
||||
self.start_idx = start_idx
|
||||
self.args = args
|
||||
self.temporal_sample = temporal_sample
|
||||
|
||||
# Initialize tokenizer
|
||||
tokenizer_path = os.path.join(args.model_path, "tokenizer")
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
cache_dir=args.cache_dir)
|
||||
|
||||
# Initialize processing stages
|
||||
self._init_stages(args, transform, transform_topcrop, tokenizer)
|
||||
|
||||
# Process metadata
|
||||
self.processed_batches = self._process_metadata()
|
||||
|
||||
def _init_stages(self, args, transform, transform_topcrop,
|
||||
tokenizer) -> None:
|
||||
"""Initialize all processing stages."""
|
||||
self.validation_stage = DataValidationStage()
|
||||
self.resolution_filter_stage = ResolutionFilterStage(
|
||||
max_height=args.max_height, max_width=args.max_width)
|
||||
self.frame_sampling_stage = FrameSamplingStage(
|
||||
num_frames=args.num_frames,
|
||||
train_fps=args.train_fps,
|
||||
speed_factor=args.speed_factor,
|
||||
video_length_tolerance_range=args.video_length_tolerance_range,
|
||||
drop_short_ratio=args.drop_short_ratio)
|
||||
self.video_transform_stage = VideoTransformStage(transform)
|
||||
self.image_transform_stage = ImageTransformStage(
|
||||
transform, transform_topcrop)
|
||||
self.text_encoding_stage = TextEncodingStage(
|
||||
tokenizer=tokenizer,
|
||||
text_max_length=args.text_max_length,
|
||||
cfg_rate=args.training_cfg_rate)
|
||||
|
||||
def _load_raw_data(self) -> List[Dict]:
|
||||
"""Load raw data from JSON files."""
|
||||
# Read folder-annotation pairs
|
||||
with open(self.data_merge_path) as f:
|
||||
folder_anno_pairs = [
|
||||
line.strip().split(",") for line in f if line.strip()
|
||||
]
|
||||
assert len(
|
||||
folder_anno_pairs) == 1, "Only support one folder-annotation pair"
|
||||
assert len(folder_anno_pairs[0]
|
||||
) == 2, "Folder-annotation pair should have two elements"
|
||||
folder, annotation_file = folder_anno_pairs[0]
|
||||
|
||||
data_items: List[Dict] = []
|
||||
with open(annotation_file) as f:
|
||||
data_items = json.load(f)
|
||||
|
||||
# Update paths with folder prefix
|
||||
for item in data_items:
|
||||
item["path"] = opj(folder, item["path"])
|
||||
|
||||
return data_items
|
||||
|
||||
def _process_metadata(self) -> List[PreprocessBatch]:
|
||||
"""Process the raw metadata through all filtering stages."""
|
||||
raw_data = self._load_raw_data()
|
||||
processed_batches = []
|
||||
|
||||
# Initialize counters
|
||||
filter_counts = {
|
||||
"validation_failed": 0,
|
||||
"resolution_failed": 0,
|
||||
"frame_sampling_failed": 0
|
||||
}
|
||||
sample_num_frames: List[int] = []
|
||||
|
||||
for item in raw_data:
|
||||
batch = PreprocessBatch(path=item["path"],
|
||||
cap=item["cap"],
|
||||
resolution=item.get("resolution"),
|
||||
fps=item.get("fps"),
|
||||
duration=item.get("duration"))
|
||||
|
||||
# Apply filtering stages
|
||||
if not self._apply_filter_stages(batch, filter_counts):
|
||||
continue
|
||||
|
||||
# Apply frame sampling processing
|
||||
batch = self.frame_sampling_stage.process(
|
||||
batch, temporal_sample_fn=self.temporal_sample)
|
||||
|
||||
processed_batches.append(batch)
|
||||
assert batch.sample_num_frames is not None
|
||||
sample_num_frames.append(batch.sample_num_frames)
|
||||
|
||||
self._log_filtering_stats(filter_counts, sample_num_frames,
|
||||
len(raw_data), len(processed_batches))
|
||||
return processed_batches
|
||||
|
||||
def _apply_filter_stages(self, batch: PreprocessBatch,
|
||||
filter_counts: Dict[str, int]) -> bool:
|
||||
"""Apply all filter stages and update counters. Returns True if batch should be kept."""
|
||||
if not self.validation_stage.should_keep(batch):
|
||||
filter_counts["validation_failed"] += 1
|
||||
return False
|
||||
|
||||
if not self.resolution_filter_stage.should_keep(batch):
|
||||
filter_counts["resolution_failed"] += 1
|
||||
return False
|
||||
|
||||
if not self.frame_sampling_stage.should_keep(batch):
|
||||
filter_counts["frame_sampling_failed"] += 1
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def _log_filtering_stats(self, filter_counts: Dict[str, int],
|
||||
sample_num_frames: List[int], before_count: int,
|
||||
after_count: int):
|
||||
"""Log filtering statistics."""
|
||||
logger.info(
|
||||
"validation_failed: %d, resolution_failed: %d, frame_sampling_failed: %d, "
|
||||
"Counter(sample_num_frames): %s, before filter: %d, after filter: %d",
|
||||
filter_counts['validation_failed'],
|
||||
filter_counts['resolution_failed'],
|
||||
filter_counts['frame_sampling_failed'], Counter(sample_num_frames),
|
||||
before_count, after_count)
|
||||
|
||||
def __iter__(self):
|
||||
"""Iterate through processed data items."""
|
||||
for idx in range(len(self.processed_batches)):
|
||||
yield self._get_item(idx)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.processed_batches)
|
||||
|
||||
def _get_item(self, idx: int) -> Dict:
|
||||
"""Get a single processed data item."""
|
||||
batch = self.processed_batches[idx]
|
||||
|
||||
# Apply transformation stages
|
||||
batch = self.video_transform_stage.process(batch)
|
||||
batch = self.image_transform_stage.process(batch)
|
||||
batch = self.text_encoding_stage.process(batch)
|
||||
|
||||
# Build result dictionary
|
||||
result = {
|
||||
"pixel_values": batch.pixel_values,
|
||||
"text": batch.text,
|
||||
"input_ids": batch.input_ids,
|
||||
"cond_mask": batch.cond_mask,
|
||||
"path": batch.path,
|
||||
}
|
||||
|
||||
# Add video-specific fields
|
||||
if batch.is_video:
|
||||
result.update({"fps": batch.fps, "duration": batch.duration})
|
||||
|
||||
return result
|
||||
|
||||
def state_dict(self) -> Dict[str, Any]:
|
||||
"""Return state dict for checkpointing."""
|
||||
return {"processed_batches": self.processed_batches}
|
||||
|
||||
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
|
||||
"""Load state dict from checkpoint."""
|
||||
self.processed_batches = state_dict["processed_batches"]
|
||||
@@ -1,352 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from collections import Counter
|
||||
from os.path import join as opj
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from fastvideo.utils.dataset_utils import DecordInit
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
|
||||
|
||||
class SingletonMeta(type):
|
||||
_instances: dict[type, 'SingletonMeta'] = {}
|
||||
|
||||
def __call__(cls, *args, **kwargs):
|
||||
if cls not in cls._instances:
|
||||
instance = super().__call__(*args, **kwargs)
|
||||
cls._instances[cls] = instance
|
||||
return cls._instances[cls]
|
||||
|
||||
|
||||
class DataSetProg(metaclass=SingletonMeta):
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.cap_list: list[dict] = []
|
||||
self.elements: list[int] = []
|
||||
self.num_workers = 1
|
||||
self.n_elements = 0
|
||||
self.worker_elements: dict[int, list[int]] = {}
|
||||
self.n_used_elements: dict[int, int] = {}
|
||||
|
||||
def set_cap_list(self, num_workers, cap_list, n_elements) -> None:
|
||||
self.num_workers = num_workers
|
||||
self.cap_list = cap_list
|
||||
self.n_elements = n_elements
|
||||
self.elements = list(range(n_elements))
|
||||
random.shuffle(self.elements)
|
||||
print(f"n_elements: {len(self.elements)}", flush=True)
|
||||
|
||||
for i in range(self.num_workers):
|
||||
self.n_used_elements[i] = 0
|
||||
per_worker = int(
|
||||
math.ceil(len(self.elements) / float(self.num_workers)))
|
||||
start = i * per_worker
|
||||
end = min(start + per_worker, len(self.elements))
|
||||
self.worker_elements[i] = self.elements[start:end]
|
||||
|
||||
def get_item(self, work_info) -> int:
|
||||
worker_id = 0 if work_info is None else work_info.id
|
||||
|
||||
idx = self.worker_elements[worker_id][
|
||||
self.n_used_elements[worker_id] %
|
||||
len(self.worker_elements[worker_id])]
|
||||
self.n_used_elements[worker_id] += 1
|
||||
return idx
|
||||
|
||||
|
||||
dataset_prog = DataSetProg()
|
||||
|
||||
|
||||
def filter_resolution(h: int,
|
||||
w: int,
|
||||
max_h_div_w_ratio: float = 17 / 16,
|
||||
min_h_div_w_ratio: float = 8 / 16) -> bool:
|
||||
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
|
||||
|
||||
|
||||
class T2V_dataset(Dataset):
|
||||
|
||||
def __init__(self,
|
||||
args,
|
||||
transform,
|
||||
temporal_sample,
|
||||
tokenizer,
|
||||
transform_topcrop,
|
||||
start_idx=0) -> None:
|
||||
self.start_idx = start_idx
|
||||
self.data = args.data_merge_path
|
||||
self.num_frames = args.num_frames
|
||||
self.train_fps = args.train_fps
|
||||
self.use_image_num = args.use_image_num
|
||||
self.transform = transform
|
||||
self.transform_topcrop = transform_topcrop
|
||||
self.temporal_sample = temporal_sample
|
||||
self.tokenizer = tokenizer
|
||||
self.text_max_length = args.text_max_length
|
||||
self.cfg = args.cfg
|
||||
self.speed_factor = args.speed_factor
|
||||
self.max_height = args.max_height
|
||||
self.max_width = args.max_width
|
||||
self.drop_short_ratio = args.drop_short_ratio
|
||||
assert self.speed_factor >= 1
|
||||
self.v_decoder = DecordInit()
|
||||
self.video_length_tolerance_range = args.video_length_tolerance_range
|
||||
self.support_Chinese = True
|
||||
if "mt5" not in args.text_encoder_name:
|
||||
self.support_Chinese = False
|
||||
|
||||
cap_list = self.get_cap_list()
|
||||
|
||||
assert len(cap_list) > 0
|
||||
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
|
||||
self.lengths = self.sample_num_frames
|
||||
|
||||
n_elements = len(cap_list)
|
||||
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list,
|
||||
n_elements)
|
||||
|
||||
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
|
||||
|
||||
def set_checkpoint(self, n_used_elements):
|
||||
for i in range(len(dataset_prog.n_used_elements)):
|
||||
dataset_prog.n_used_elements[i] = n_used_elements
|
||||
|
||||
def __len__(self):
|
||||
return dataset_prog.n_elements
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
data = self.get_data(idx)
|
||||
return data
|
||||
|
||||
def get_data(self, idx) -> dict:
|
||||
path = dataset_prog.cap_list[idx]["path"]
|
||||
if path.endswith(".mp4"):
|
||||
return self.get_video(idx)
|
||||
else:
|
||||
return self.get_image(idx)
|
||||
|
||||
def get_video(self, idx) -> dict:
|
||||
video_path = dataset_prog.cap_list[idx]["path"]
|
||||
assert os.path.exists(video_path), f"file {video_path} do not exist!"
|
||||
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
|
||||
|
||||
torchvision_video, _, metadata = torchvision.io.read_video(
|
||||
video_path, output_format="TCHW")
|
||||
video = torchvision_video[frame_indices]
|
||||
video = self.transform(video)
|
||||
video = rearrange(video, "t c h w -> c t h w")
|
||||
video = video.to(torch.uint8)
|
||||
assert video.dtype == torch.uint8
|
||||
|
||||
h, w = video.shape[-2:]
|
||||
assert (
|
||||
h / w <= 17 / 16 and h / w >= 8 / 16
|
||||
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
|
||||
|
||||
video = video.float() / 127.5 - 1.0
|
||||
|
||||
text = dataset_prog.cap_list[idx]["cap"]
|
||||
if not isinstance(text, list):
|
||||
text = [text]
|
||||
text = [random.choice(text)]
|
||||
|
||||
text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"]
|
||||
cond_mask = text_tokens_and_mask["attention_mask"]
|
||||
return dict(pixel_values=video,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=video_path,
|
||||
fps=dataset_prog.cap_list[idx]["fps"],
|
||||
duration=dataset_prog.cap_list[idx]["duration"])
|
||||
|
||||
def get_image(self, idx) -> dict:
|
||||
image_data = dataset_prog.cap_list[
|
||||
idx] # [{'path': path, 'cap': cap}, ...]
|
||||
|
||||
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
|
||||
image = torch.from_numpy(np.array(image)) # [h, w, c]
|
||||
image = rearrange(image, "h w c -> c h w").unsqueeze(0) # [1 c h w]
|
||||
# for i in image:
|
||||
# h, w = i.shape[-2:]
|
||||
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
|
||||
|
||||
image = (self.transform_topcrop(image) if "human_images"
|
||||
in image_data["path"] else self.transform(image)
|
||||
) # [1 C H W] -> num_img [1 C H W]
|
||||
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
|
||||
|
||||
image = image.float() / 127.5 - 1.0
|
||||
|
||||
caps: list[str] = (image_data["cap"] if isinstance(
|
||||
image_data["cap"], list) else [image_data["cap"]])
|
||||
caps = [random.choice(caps)]
|
||||
text = caps
|
||||
input_ids, cond_mask = [], []
|
||||
single_text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
single_text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"] # 1, l
|
||||
cond_mask = text_tokens_and_mask["attention_mask"] # 1, l
|
||||
return dict(
|
||||
pixel_values=image,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=image_data["path"],
|
||||
)
|
||||
|
||||
def define_frame_index(self, cap_list) -> tuple[list[dict], list[int]]:
|
||||
new_cap_list = []
|
||||
sample_num_frames = []
|
||||
cnt_too_long = 0
|
||||
cnt_too_short = 0
|
||||
cnt_no_cap = 0
|
||||
cnt_no_resolution = 0
|
||||
cnt_resolution_mismatch = 0
|
||||
cnt_movie = 0
|
||||
cnt_img = 0
|
||||
for i in cap_list:
|
||||
path = i["path"]
|
||||
cap = i.get("cap", None)
|
||||
# ======no caption=====
|
||||
if cap is None:
|
||||
cnt_no_cap += 1
|
||||
continue
|
||||
if path.endswith(".mp4"):
|
||||
# ======no fps and duration=====
|
||||
duration = i.get("duration", None)
|
||||
fps = i.get("fps", None)
|
||||
if fps is None or duration is None:
|
||||
continue
|
||||
|
||||
# ======resolution mismatch=====
|
||||
resolution = i.get("resolution", None)
|
||||
if resolution is None:
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
else:
|
||||
if (resolution.get("height", None) is None
|
||||
or resolution.get("width", None) is None):
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
height, width = i["resolution"]["height"], i["resolution"][
|
||||
"width"]
|
||||
aspect = self.max_height / self.max_width
|
||||
hw_aspect_thr = 1.5
|
||||
is_pick = filter_resolution(
|
||||
height,
|
||||
width,
|
||||
max_h_div_w_ratio=hw_aspect_thr * aspect,
|
||||
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
|
||||
)
|
||||
if not is_pick:
|
||||
print("resolution mismatch")
|
||||
cnt_resolution_mismatch += 1
|
||||
continue
|
||||
|
||||
# if path == 'finetrainers/3dgs-dissolve/videos/1.mp4':
|
||||
# from IPython import embed; embed()
|
||||
i["num_frames"] = math.ceil(fps * duration)
|
||||
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
|
||||
if i["num_frames"] / fps > self.video_length_tolerance_range * (
|
||||
self.num_frames / self.train_fps * self.speed_factor
|
||||
): # too long video is not suitable for this training stage (self.num_frames)
|
||||
cnt_too_long += 1
|
||||
continue
|
||||
|
||||
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
|
||||
frame_interval = fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, i["num_frames"],
|
||||
frame_interval).astype(int)
|
||||
|
||||
# comment out it to enable dynamic frames training
|
||||
if (len(frame_indices) < self.num_frames
|
||||
and random.random() < self.drop_short_ratio):
|
||||
cnt_too_short += 1
|
||||
continue
|
||||
|
||||
# too long video will be temporal-crop randomly
|
||||
if len(frame_indices) > self.num_frames:
|
||||
begin_index, end_index = self.temporal_sample(
|
||||
len(frame_indices))
|
||||
frame_indices = frame_indices[begin_index:end_index]
|
||||
# frame_indices = frame_indices[:self.num_frames] # head crop
|
||||
i["sample_frame_index"] = frame_indices.tolist()
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = len(
|
||||
i["sample_frame_index"]
|
||||
) # will use in dataloader(group sampler)
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
elif path.endswith(".jpg"): # image
|
||||
cnt_img += 1
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = 1
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
else:
|
||||
raise NameError(
|
||||
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
|
||||
)
|
||||
# import ipdb;ipdb.set_trace()
|
||||
main_print(
|
||||
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
|
||||
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
|
||||
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
|
||||
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}"
|
||||
)
|
||||
return new_cap_list, sample_num_frames
|
||||
|
||||
def decord_read(self, path, frame_indices) -> torch.Tensor:
|
||||
decord_vr = self.v_decoder(path)
|
||||
video_data = decord_vr.get_batch(frame_indices).asnumpy()
|
||||
video_data = torch.from_numpy(video_data)
|
||||
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
|
||||
return video_data
|
||||
|
||||
def read_jsons(self, data) -> list[dict]:
|
||||
cap_lists = []
|
||||
with open(data) as f:
|
||||
folder_anno = [
|
||||
i.strip().split(",") for i in f.readlines()
|
||||
if len(i.strip()) > 0
|
||||
]
|
||||
print(folder_anno)
|
||||
for folder, anno in folder_anno:
|
||||
with open(anno) as f:
|
||||
sub_list = json.load(f)
|
||||
for i in range(len(sub_list)):
|
||||
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
|
||||
cap_lists += sub_list
|
||||
return cap_lists
|
||||
|
||||
def get_cap_list(self) -> list:
|
||||
cap_lists = self.read_jsons(self.data)[self.start_idx:]
|
||||
return cap_lists
|
||||
@@ -1,4 +1,5 @@
|
||||
from typing import Any, Dict, List
|
||||
import random
|
||||
from typing import Any, Dict, List, cast
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -20,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,7 +43,10 @@ def get_torch_tensors_from_row_dict(row_dict, keys) -> Dict[str, Any]:
|
||||
bytes = row_dict[f"{key}_bytes"]
|
||||
|
||||
# TODO (peiyuan): read precision
|
||||
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
|
||||
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
|
||||
@@ -53,8 +57,11 @@ def get_torch_tensors_from_row_dict(row_dict, keys) -> Dict[str, Any]:
|
||||
|
||||
|
||||
def collate_latents_embs_masks(
|
||||
batch_to_process, text_padding_length,
|
||||
keys) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, List[str]]:
|
||||
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 = []
|
||||
@@ -63,7 +70,7 @@ def collate_latents_embs_masks(
|
||||
# Process each row individually
|
||||
for i, row in enumerate(batch_to_process):
|
||||
# Get tensors from row
|
||||
data = get_torch_tensors_from_row_dict(row, keys)
|
||||
data = get_torch_tensors_from_row_dict(row, keys, cfg_rate)
|
||||
latents, emb = data["vae_latent"], data["text_embedding"]
|
||||
|
||||
padded_emb, mask = pad(emb, text_padding_length)
|
||||
@@ -83,3 +90,131 @@ def collate_latents_embs_masks(
|
||||
all_masks = torch.stack(all_masks)
|
||||
|
||||
return all_latents, all_embs, all_masks, caption_text
|
||||
|
||||
|
||||
def collate_rows_from_parquet_schema(rows,
|
||||
parquet_schema,
|
||||
text_padding_length,
|
||||
cfg_rate=0.0) -> Dict[str, Any]:
|
||||
"""
|
||||
Collate rows from parquet files based on the provided schema.
|
||||
Dynamically processes tensor fields based on schema and returns batched data.
|
||||
|
||||
Args:
|
||||
rows: List of row dictionaries from parquet files
|
||||
parquet_schema: PyArrow schema defining the structure of the data
|
||||
|
||||
Returns:
|
||||
Dict containing batched tensors and metadata
|
||||
"""
|
||||
if not rows:
|
||||
return cast(Dict[str, Any], {})
|
||||
|
||||
# Initialize containers for different data types
|
||||
batch_data: Dict[str, Any] = {}
|
||||
|
||||
# Get tensor and metadata field names from schema (fields ending with '_bytes')
|
||||
tensor_fields = []
|
||||
metadata_fields = []
|
||||
for field in parquet_schema.names:
|
||||
if field.endswith('_bytes'):
|
||||
shape_field = field.replace('_bytes', '_shape')
|
||||
dtype_field = field.replace('_bytes', '_dtype')
|
||||
tensor_name = field.replace('_bytes', '')
|
||||
tensor_fields.append(tensor_name)
|
||||
assert shape_field in parquet_schema.names, f"Shape field {shape_field} not found in schema for field {field}. Currently we only support *_bytes fields for tensors."
|
||||
assert dtype_field in parquet_schema.names, f"Dtype field {dtype_field} not found in schema for field {field}. Currently we only support *_bytes fields for tensors."
|
||||
elif not field.endswith('_shape') and not field.endswith('_dtype'):
|
||||
# Only add actual metadata fields, not the shape/dtype helper fields
|
||||
metadata_fields.append(field)
|
||||
|
||||
# Process each tensor field
|
||||
for tensor_name in tensor_fields:
|
||||
tensor_list = []
|
||||
|
||||
for row in rows:
|
||||
# Get tensor data from row using the existing helper function pattern
|
||||
shape_key = f"{tensor_name}_shape"
|
||||
bytes_key = f"{tensor_name}_bytes"
|
||||
|
||||
if shape_key in row and bytes_key in row:
|
||||
shape = row[shape_key]
|
||||
bytes_data = row[bytes_key]
|
||||
|
||||
if len(bytes_data) == 0:
|
||||
tensor = torch.zeros(0, dtype=torch.bfloat16)
|
||||
else:
|
||||
# Convert bytes to tensor using float32 as default
|
||||
if tensor_name == 'text_embedding' and random.random(
|
||||
) < cfg_rate:
|
||||
data = np.zeros((512, 4096), dtype=np.float32)
|
||||
else:
|
||||
data = np.frombuffer(
|
||||
bytes_data, dtype=np.float32).reshape(shape).copy()
|
||||
tensor = torch.from_numpy(data)
|
||||
# if len(data.shape) == 3:
|
||||
# B, L, D = tensor.shape
|
||||
# assert B == 1, "Batch size must be 1"
|
||||
# tensor = tensor.squeeze(0)
|
||||
|
||||
tensor_list.append(tensor)
|
||||
else:
|
||||
# Handle missing tensor data
|
||||
tensor_list.append(torch.zeros(0, dtype=torch.bfloat16))
|
||||
|
||||
# Stack tensors with special handling for text embeddings
|
||||
if tensor_name == 'text_embedding':
|
||||
# Handle text embeddings with padding
|
||||
padded_tensors = []
|
||||
attention_masks = []
|
||||
|
||||
for tensor in tensor_list:
|
||||
if tensor.numel() > 0:
|
||||
padded_tensor, mask = pad(tensor, text_padding_length)
|
||||
padded_tensors.append(padded_tensor)
|
||||
attention_masks.append(mask)
|
||||
else:
|
||||
# Handle empty embeddings - assume default embedding dimension
|
||||
padded_tensors.append(
|
||||
torch.zeros(text_padding_length,
|
||||
768,
|
||||
dtype=torch.bfloat16))
|
||||
attention_masks.append(torch.zeros(text_padding_length))
|
||||
|
||||
batch_data[tensor_name] = torch.stack(padded_tensors)
|
||||
batch_data['text_attention_mask'] = torch.stack(attention_masks)
|
||||
else:
|
||||
# Stack all tensors to preserve batch consistency
|
||||
# Don't filter out None or empty tensors as this breaks batch sizing
|
||||
try:
|
||||
batch_data[tensor_name] = torch.stack(tensor_list)
|
||||
except ValueError as e:
|
||||
shapes = [
|
||||
t.shape
|
||||
if t is not None and hasattr(t, 'shape') else 'None/Invalid'
|
||||
for t in tensor_list
|
||||
]
|
||||
raise ValueError(
|
||||
f"Failed to stack tensors for field '{tensor_name}'. "
|
||||
f"Tensor shapes: {shapes}. "
|
||||
f"All tensors in a batch must have compatible shapes. "
|
||||
f"Original error: {e}") from e
|
||||
|
||||
# Process metadata fields into info_list
|
||||
info_list = []
|
||||
for row in rows:
|
||||
info = {}
|
||||
for field in metadata_fields:
|
||||
info[field] = row.get(field, "")
|
||||
|
||||
# Add prompt field for backward compatibility
|
||||
info["prompt"] = info.get("caption", "")
|
||||
info_list.append(info)
|
||||
|
||||
batch_data['info_list'] = info_list
|
||||
|
||||
# Add caption_text for backward compatibility
|
||||
if info_list and 'caption' in info_list[0]:
|
||||
batch_data['caption_text'] = [info['caption'] for info in info_list]
|
||||
|
||||
return batch_data
|
||||
|
||||
@@ -0,0 +1,103 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from: https://github.com/a-r-r-o-w/finetrainers/blob/main/finetrainers/data/dataset.py
|
||||
import os
|
||||
import pathlib
|
||||
|
||||
import datasets
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vision_utils import load_image, load_video
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ValidationDataset(torch.utils.data.IterableDataset):
|
||||
|
||||
def __init__(self, filename: str):
|
||||
super().__init__()
|
||||
|
||||
self.filename = pathlib.Path(filename)
|
||||
# get directory of filename
|
||||
self.dir = os.path.abspath(self.filename.parent)
|
||||
|
||||
if not self.filename.exists():
|
||||
raise FileNotFoundError(
|
||||
f"File {self.filename.as_posix()} does not exist")
|
||||
|
||||
if self.filename.suffix == ".csv":
|
||||
data = datasets.load_dataset("csv",
|
||||
data_files=self.filename.as_posix(),
|
||||
split="train")
|
||||
elif self.filename.suffix == ".json":
|
||||
data = datasets.load_dataset("json",
|
||||
data_files=self.filename.as_posix(),
|
||||
split="train",
|
||||
field="data")
|
||||
elif self.filename.suffix == ".parquet":
|
||||
data = datasets.load_dataset("parquet",
|
||||
data_files=self.filename.as_posix(),
|
||||
split="train")
|
||||
elif self.filename.suffix == ".arrow":
|
||||
data = datasets.load_dataset("arrow",
|
||||
data_files=self.filename.as_posix(),
|
||||
split="train")
|
||||
else:
|
||||
_SUPPORTED_FILE_FORMATS = [".csv", ".json", ".parquet", ".arrow"]
|
||||
raise ValueError(
|
||||
f"Unsupported file format {self.filename.suffix} for validation dataset. Supported formats are: {_SUPPORTED_FILE_FORMATS}"
|
||||
)
|
||||
|
||||
self._data = data.to_iterable_dataset()
|
||||
|
||||
def __iter__(self):
|
||||
for sample in self._data:
|
||||
# For consistency reasons, we mandate that "caption" is always present in the validation dataset.
|
||||
# However, since the model specifications use "prompt", we create an alias here.
|
||||
sample["prompt"] = sample["caption"]
|
||||
|
||||
# Load image or video if the path is provided
|
||||
# TODO(aryan): need to handle custom columns here for control conditions
|
||||
sample["image"] = None
|
||||
sample["video"] = None
|
||||
|
||||
if sample.get("image_path", None) is not None:
|
||||
image_path = sample["image_path"]
|
||||
image_path = os.path.join(self.dir, image_path)
|
||||
if not pathlib.Path(image_path).is_file(
|
||||
) and not image_path.startswith("http"):
|
||||
logger.warning("Image file %s does not exist.", image_path)
|
||||
else:
|
||||
sample["image"] = load_image(image_path)
|
||||
|
||||
if sample.get("video_path", None) is not None:
|
||||
video_path = sample["video_path"]
|
||||
video_path = os.path.join(self.dir, video_path)
|
||||
if not pathlib.Path(video_path).is_file(
|
||||
) and not video_path.startswith("http"):
|
||||
logger.warning("Video file %s does not exist.", video_path)
|
||||
else:
|
||||
sample["video"] = load_video(video_path)
|
||||
|
||||
if sample.get("control_image_path", None) is not None:
|
||||
control_image_path = sample["control_image_path"]
|
||||
control_image_path = os.path.join(self.dir, control_image_path)
|
||||
if not pathlib.Path(control_image_path).is_file(
|
||||
) and not control_image_path.startswith("http"):
|
||||
logger.warning("Control Image file %s does not exist.",
|
||||
control_image_path)
|
||||
else:
|
||||
sample["control_image"] = load_image(control_image_path)
|
||||
|
||||
if sample.get("control_video_path", None) is not None:
|
||||
control_video_path = sample["control_video_path"]
|
||||
control_video_path = os.path.join(self.dir, control_video_path)
|
||||
if not pathlib.Path(control_video_path).is_file(
|
||||
) and not control_video_path.startswith("http"):
|
||||
logger.warning("Control Video file %s does not exist.",
|
||||
control_video_path)
|
||||
else:
|
||||
sample["control_video"] = load_video(control_video_path)
|
||||
|
||||
sample = {k: v for k, v in sample.items() if v is not None}
|
||||
yield sample
|
||||
@@ -8,7 +8,7 @@ from contextlib import contextmanager
|
||||
from dataclasses import field
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig, STA_Mode
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser, StoreBoolean
|
||||
|
||||
@@ -63,7 +63,7 @@ class FastVideoArgs:
|
||||
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
STA_mode: Optional[str] = None
|
||||
STA_mode: STA_Mode = STA_Mode.STA_INFERENCE
|
||||
skip_time_steps: int = 15
|
||||
|
||||
# Compilation
|
||||
@@ -74,6 +74,9 @@ class FastVideoArgs:
|
||||
# VSA parameters
|
||||
VSA_sparsity: float = 0.0 # inference/validation sparsity
|
||||
|
||||
# Stage verification
|
||||
enable_stage_verification: bool = True
|
||||
|
||||
@property
|
||||
def training_mode(self) -> bool:
|
||||
return not self.inference_mode
|
||||
@@ -178,12 +181,10 @@ class FastVideoArgs:
|
||||
parser.add_argument(
|
||||
"--STA-mode",
|
||||
type=str,
|
||||
default=FastVideoArgs.STA_mode,
|
||||
choices=[
|
||||
"STA_inference", "STA_searching", "STA_tuning",
|
||||
"STA_tuning_cfg", None
|
||||
],
|
||||
help="STA mode",
|
||||
default=FastVideoArgs.STA_mode.value,
|
||||
choices=[mode.value for mode in STA_Mode],
|
||||
help=
|
||||
"STA mode contains STA_inference, STA_searching, STA_tuning, STA_tuning_cfg, None",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--skip-time-steps",
|
||||
@@ -231,6 +232,14 @@ class FastVideoArgs:
|
||||
help="Validation sparsity for VSA",
|
||||
)
|
||||
|
||||
# Stage verification
|
||||
parser.add_argument(
|
||||
"--enable-stage-verification",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.enable_stage_verification,
|
||||
help="Enable input/output verification for pipeline stages",
|
||||
)
|
||||
|
||||
# Add pipeline configuration arguments
|
||||
PipelineConfig.add_cli_args(parser)
|
||||
|
||||
@@ -375,11 +384,12 @@ 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
|
||||
validation_prompt_dir: str = ""
|
||||
validation_dataset_file: str = ""
|
||||
validation_preprocessed_path: str = ""
|
||||
validation_sampling_steps: str = ""
|
||||
validation_guidance_scale: str = ""
|
||||
validation_steps: float = 0.0
|
||||
@@ -403,7 +413,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
lr_scheduler: str = "constant"
|
||||
lr_warmup_steps: int = 0
|
||||
max_grad_norm: float = 0.0
|
||||
gradient_checkpointing: bool = False
|
||||
enable_gradient_checkpointing_type: Optional[str] = None
|
||||
selective_checkpointing: float = 0.0
|
||||
allow_tf32: bool = False
|
||||
mixed_precision: str = ""
|
||||
@@ -518,7 +528,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(
|
||||
@@ -527,9 +537,12 @@ class TrainingArgs(FastVideoArgs):
|
||||
help="Whether to precondition the outputs of the model")
|
||||
|
||||
# Validation and logging
|
||||
parser.add_argument("--validation-prompt-dir",
|
||||
parser.add_argument("--validation-dataset-file",
|
||||
type=str,
|
||||
help="Directory containing validation prompts")
|
||||
help="Path to unprocessed validation dataset")
|
||||
parser.add_argument("--validation-preprocessed-path",
|
||||
type=str,
|
||||
help="Path to processed validation dataset")
|
||||
parser.add_argument("--validation-sampling-steps",
|
||||
type=str,
|
||||
help="Validation sampling steps")
|
||||
@@ -599,9 +612,11 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--max-grad-norm",
|
||||
type=float,
|
||||
help="Maximum gradient norm")
|
||||
parser.add_argument("--gradient-checkpointing",
|
||||
action=StoreBoolean,
|
||||
help="Whether to use gradient checkpointing")
|
||||
parser.add_argument("--enable-gradient-checkpointing-type",
|
||||
type=str,
|
||||
choices=["full", "ops", "block_skip"],
|
||||
default=None,
|
||||
help="Gradient checkpointing type")
|
||||
parser.add_argument("--selective-checkpointing",
|
||||
type=float,
|
||||
help="Selective checkpointing threshold")
|
||||
|
||||
@@ -5,15 +5,16 @@ import time
|
||||
from collections import defaultdict
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
# if TYPE_CHECKING:
|
||||
from fastvideo.v1.attention import AttentionMetadata
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.v1.attention import AttentionMetadata
|
||||
from fastvideo.v1.pipelines import ForwardBatch
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -36,13 +37,13 @@ class ForwardContext:
|
||||
# attn_layers: Dict[str, Any]
|
||||
# TODO: extend to support per-layer dynamic forward context
|
||||
attn_metadata: "AttentionMetadata" # set dynamically for each forward pass
|
||||
forward_batch: Optional[ForwardBatch] = None
|
||||
forward_batch: Optional["ForwardBatch"] = None
|
||||
|
||||
|
||||
_forward_context: Optional[ForwardContext] = None
|
||||
_forward_context: Optional["ForwardContext"] = None
|
||||
|
||||
|
||||
def get_forward_context() -> ForwardContext:
|
||||
def get_forward_context() -> "ForwardContext":
|
||||
"""Get the current forward context."""
|
||||
assert _forward_context is not None, (
|
||||
"Forward context is not set. "
|
||||
@@ -54,7 +55,7 @@ def get_forward_context() -> ForwardContext:
|
||||
@contextmanager
|
||||
def set_forward_context(current_timestep,
|
||||
attn_metadata,
|
||||
forward_batch: Optional[ForwardBatch] = None,
|
||||
forward_batch: Optional["ForwardBatch"] = None,
|
||||
fastvideo_args: Optional[FastVideoArgs] = None):
|
||||
"""A context manager that stores the current forward context,
|
||||
can be attention metadata, etc.
|
||||
|
||||
@@ -5,6 +5,7 @@ from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.v1.layers.custom_op import CustomOp
|
||||
|
||||
@@ -37,6 +38,12 @@ class RMSNorm(CustomOp):
|
||||
if self.has_weight:
|
||||
self.weight = nn.Parameter(self.weight)
|
||||
|
||||
# if we do fully_shard(model.layer_norm), and we call layer_form.forward_native(input) instead of layer_norm(input),
|
||||
# we need to call model.layer_norm.register_fsdp_forward_method(model, "forward_native") to make sure fsdp2 hooks are triggered
|
||||
# for mixed precision and cpu offloading
|
||||
|
||||
# the even better way might be fully_shard(model.layer_norm, mp_policy=, cpu_offloading=), and call model.layer_norm(input). everything should work out of the box
|
||||
# because fsdp2 hooks will be triggered with model.layer_norm.__call__
|
||||
def forward_native(
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
@@ -95,6 +102,22 @@ class ScaleResidual(nn.Module):
|
||||
return residual + x * gate
|
||||
|
||||
|
||||
# adapted from Diffusers: https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/normalization.py
|
||||
# NOTE(will): Needed to match behavior of diffusers and wan2.1 even while using
|
||||
# FSDP's MixedPrecisionPolicy
|
||||
class FP32LayerNorm(nn.LayerNorm):
|
||||
|
||||
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
|
||||
origin_dtype = inputs.dtype
|
||||
return F.layer_norm(
|
||||
inputs.float(),
|
||||
self.normalized_shape,
|
||||
self.weight.float() if self.weight is not None else None,
|
||||
self.bias.float() if self.bias is not None else None,
|
||||
self.eps,
|
||||
).to(origin_dtype)
|
||||
|
||||
|
||||
class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
"""
|
||||
Fused operation that combines:
|
||||
@@ -112,6 +135,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
eps: float = 1e-6,
|
||||
elementwise_affine: bool = False,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
compute_dtype: torch.dtype | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
@@ -121,10 +145,15 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
elif norm_type == "layer":
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
if compute_dtype == torch.float32:
|
||||
self.norm = FP32LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps)
|
||||
else:
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
else:
|
||||
raise NotImplementedError(f"Norm type {norm_type} not implemented")
|
||||
|
||||
@@ -163,18 +192,25 @@ class LayerNormScaleShift(nn.Module):
|
||||
eps: float = 1e-6,
|
||||
elementwise_affine: bool = False,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
compute_dtype: torch.dtype | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.compute_dtype = compute_dtype
|
||||
if norm_type == "rms":
|
||||
self.norm = RMSNorm(hidden_size,
|
||||
has_weight=elementwise_affine,
|
||||
eps=eps)
|
||||
elif norm_type == "layer":
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
if self.compute_dtype == torch.float32:
|
||||
self.norm = FP32LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps)
|
||||
else:
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
else:
|
||||
raise NotImplementedError(f"Norm type {norm_type} not implemented")
|
||||
|
||||
@@ -182,4 +218,7 @@ class LayerNormScaleShift(nn.Module):
|
||||
scale: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply ln followed by scale and shift in a single fused operation."""
|
||||
normalized = self.norm(x)
|
||||
return normalized * (1.0 + scale) + shift
|
||||
if self.compute_dtype == torch.float32:
|
||||
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
|
||||
else:
|
||||
return normalized * (1.0 + scale) + shift
|
||||
|
||||
@@ -71,15 +71,16 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
f"cuda:{torch.cuda.current_device()}").full_tensor()
|
||||
data += (self.slice_lora_b_weights(self.lora_B)
|
||||
@ self.slice_lora_a_weights(self.lora_A)).to(data)
|
||||
self.base_layer.weight.data = distribute_tensor(
|
||||
data, mesh, placements=placements).to(current_device)
|
||||
self.base_layer.weight = nn.Parameter(
|
||||
distribute_tensor(data, mesh,
|
||||
placements=placements).to(current_device))
|
||||
else:
|
||||
current_device = self.base_layer.weight.data.device
|
||||
data = self.base_layer.weight.data.to(
|
||||
data = self.base_layer.weight.to(
|
||||
f"cuda:{torch.cuda.current_device()}")
|
||||
data += \
|
||||
(self.slice_lora_b_weights(self.lora_B) @ self.slice_lora_a_weights(self.lora_A)).to(data)
|
||||
self.base_layer.weight.data = data.to(current_device)
|
||||
self.base_layer.weight = nn.Parameter(data.to(current_device))
|
||||
self.merged = True
|
||||
|
||||
@torch.no_grad()
|
||||
@@ -106,8 +107,8 @@ class BaseLayerWithLoRA(nn.Module):
|
||||
f"cuda:{torch.cuda.current_device()}").full_tensor()
|
||||
data -= self.slice_lora_b_weights(
|
||||
self.lora_B) @ self.slice_lora_a_weights(self.lora_A)
|
||||
self.base_layer.weight.data = distribute_tensor(
|
||||
data, mesh, placements=placement).to(device)
|
||||
self.base_layer.weight = nn.Parameter(
|
||||
distribute_tensor(data, mesh, placements=placement).to(device))
|
||||
else:
|
||||
self.base_layer.weight.data -= \
|
||||
self.slice_lora_b_weights(self.lora_B) @\
|
||||
|
||||
@@ -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]):
|
||||
|
||||
@@ -460,6 +460,8 @@ class StepVideoModel(BaseDiT):
|
||||
# lambda n, m: "pos_embed" in n # If needed for the patch embedding.
|
||||
]
|
||||
_param_names_mapping = StepVideoConfig()._param_names_mapping
|
||||
_reverse_param_names_mapping = StepVideoConfig(
|
||||
)._reverse_param_names_mapping
|
||||
_lora_param_names_mapping = StepVideoConfig()._lora_param_names_mapping
|
||||
_supported_attention_backends = StepVideoConfig(
|
||||
)._supported_attention_backends
|
||||
|
||||
@@ -14,8 +14,8 @@ from fastvideo.v1.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.v1.configs.sample.wan import WanTeaCacheParams
|
||||
from fastvideo.v1.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.v1.forward_context import get_forward_context
|
||||
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, RMSNorm,
|
||||
ScaleResidual,
|
||||
from fastvideo.v1.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
|
||||
RMSNorm, ScaleResidual,
|
||||
ScaleResidualLayerNormScaleShift)
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
# from torch.nn import RMSNorm
|
||||
@@ -34,9 +34,9 @@ class WanImageEmbedding(torch.nn.Module):
|
||||
def __init__(self, in_features: int, out_features: int):
|
||||
super().__init__()
|
||||
|
||||
self.norm1 = nn.LayerNorm(in_features)
|
||||
self.norm1 = FP32LayerNorm(in_features)
|
||||
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
|
||||
self.norm2 = nn.LayerNorm(out_features)
|
||||
self.norm2 = FP32LayerNorm(out_features)
|
||||
|
||||
def forward(self,
|
||||
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
@@ -232,7 +232,7 @@ class WanTransformerBlock(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -263,7 +263,8 @@ class WanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
@@ -283,7 +284,8 @@ class WanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -375,7 +377,7 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -407,7 +409,8 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
@@ -427,7 +430,8 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -514,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,
|
||||
@@ -564,7 +569,8 @@ class WanTransformer3DModel(CachableDiT):
|
||||
norm_type="layer",
|
||||
eps=config.eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
@@ -572,6 +578,17 @@ class WanTransformer3DModel(CachableDiT):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
# For type checking
|
||||
self.previous_e0_even = None
|
||||
self.previous_e0_odd = None
|
||||
self.previous_residual_even = None
|
||||
self.previous_residual_odd = None
|
||||
self.is_even = True
|
||||
self.should_calc_even = True
|
||||
self.should_calc_odd = True
|
||||
self.accumulated_rel_l1_distance_even = 0
|
||||
self.accumulated_rel_l1_distance_odd = 0
|
||||
self.cnt = 0
|
||||
self.__post_init__()
|
||||
|
||||
def forward(self,
|
||||
@@ -658,7 +675,7 @@ class WanTransformer3DModel(CachableDiT):
|
||||
# 5. Output norm, projection & unpatchify
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
|
||||
dim=1)
|
||||
hidden_states = self.norm_out(hidden_states.float(), shift, scale)
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
|
||||
@@ -222,10 +222,14 @@ 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
|
||||
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 +264,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",
|
||||
|
||||
@@ -51,11 +51,6 @@ def auto_attributes(init_func):
|
||||
return wrapper
|
||||
|
||||
|
||||
def set_random_seed(seed: int) -> None:
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
current_platform.seed_everything(seed)
|
||||
|
||||
|
||||
def set_weight_attrs(
|
||||
weight: torch.Tensor,
|
||||
weight_attrs: Optional[Dict[str, Any]],
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from typing import Callable, List, Optional, Tuple, Union
|
||||
from urllib.parse import unquote, urlparse
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
import PIL.ImageOps
|
||||
@@ -86,6 +89,7 @@ def normalize(
|
||||
return 2.0 * images - 1.0
|
||||
|
||||
|
||||
# adapted from diffusers.utils import load_image
|
||||
def load_image(
|
||||
image: Union[str, PIL.Image.Image],
|
||||
convert_method: Optional[Callable[[PIL.Image.Image],
|
||||
@@ -131,6 +135,85 @@ def load_image(
|
||||
return image
|
||||
|
||||
|
||||
# adapted from diffusers.utils import load_video
|
||||
def load_video(
|
||||
video: str,
|
||||
convert_method: Optional[Callable[[List[PIL.Image.Image]],
|
||||
List[PIL.Image.Image]]] = None,
|
||||
) -> List[PIL.Image.Image]:
|
||||
"""
|
||||
Loads `video` to a list of PIL Image.
|
||||
Args:
|
||||
video (`str`):
|
||||
A URL or Path to a video to convert to a list of PIL Image format.
|
||||
convert_method (Callable[[List[PIL.Image.Image]], List[PIL.Image.Image]], *optional*):
|
||||
A conversion method to apply to the video after loading it. When set to `None` the images will be converted
|
||||
to "RGB".
|
||||
Returns:
|
||||
`List[PIL.Image.Image]`:
|
||||
The video as a list of PIL images.
|
||||
"""
|
||||
is_url = video.startswith("http://") or video.startswith("https://")
|
||||
is_file = os.path.isfile(video)
|
||||
was_tempfile_created = False
|
||||
|
||||
if not (is_url or is_file):
|
||||
raise ValueError(
|
||||
f"Incorrect path or URL. URLs must start with `http://` or `https://`, and {video} is not a valid path."
|
||||
)
|
||||
|
||||
if is_url:
|
||||
response = requests.get(video, stream=True)
|
||||
if response.status_code != 200:
|
||||
raise ValueError(
|
||||
f"Failed to download video. Status code: {response.status_code}"
|
||||
)
|
||||
|
||||
parsed_url = urlparse(video)
|
||||
file_name = os.path.basename(unquote(parsed_url.path))
|
||||
|
||||
suffix = os.path.splitext(file_name)[1] or ".mp4"
|
||||
with tempfile.NamedTemporaryFile(suffix=suffix,
|
||||
delete=False) as temp_file:
|
||||
video_path = temp_file.name
|
||||
video_data = response.iter_content(chunk_size=8192)
|
||||
for chunk in video_data:
|
||||
temp_file.write(chunk)
|
||||
|
||||
video = video_path
|
||||
|
||||
pil_images = []
|
||||
if video.endswith(".gif"):
|
||||
gif = PIL.Image.open(video)
|
||||
try:
|
||||
while True:
|
||||
pil_images.append(gif.copy())
|
||||
gif.seek(gif.tell() + 1)
|
||||
except EOFError:
|
||||
pass
|
||||
|
||||
else:
|
||||
try:
|
||||
imageio.plugins.ffmpeg.get_exe()
|
||||
except AttributeError:
|
||||
raise AttributeError(
|
||||
"`Unable to find an ffmpeg installation on your machine. Please install via `pip install imageio-ffmpeg"
|
||||
) from None
|
||||
|
||||
with imageio.get_reader(video) as reader:
|
||||
# Read all frames
|
||||
for frame in reader:
|
||||
pil_images.append(PIL.Image.fromarray(frame))
|
||||
|
||||
if was_tempfile_created:
|
||||
os.remove(video_path)
|
||||
|
||||
if convert_method is not None:
|
||||
pil_images = convert_method(pil_images)
|
||||
|
||||
return pil_images
|
||||
|
||||
|
||||
def get_default_height_width(
|
||||
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
|
||||
vae_scale_factor: int,
|
||||
|
||||
@@ -11,7 +11,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.lora_pipeline import LoRAPipeline
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
|
||||
TrainingBatch)
|
||||
from fastvideo.v1.pipelines.pipeline_registry import PipelineRegistry
|
||||
from fastvideo.v1.utils import (maybe_download_model,
|
||||
verify_model_config_and_directory)
|
||||
@@ -63,4 +64,5 @@ __all__ = [
|
||||
"PipelineRegistry",
|
||||
"ForwardBatch",
|
||||
"LoRAPipeline",
|
||||
"TrainingBatch",
|
||||
]
|
||||
|
||||
@@ -298,3 +298,7 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
# Return the output
|
||||
return batch
|
||||
|
||||
def train(self) -> None:
|
||||
raise NotImplementedError(
|
||||
"if training_mode is True, the pipeline must implement this method")
|
||||
|
||||
@@ -11,8 +11,10 @@ import pprint
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import PIL.Image
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.attention import AttentionMetadata
|
||||
from fastvideo.v1.configs.sample.teacache import (TeaCacheParams,
|
||||
WanTeaCacheParams)
|
||||
|
||||
@@ -36,6 +38,8 @@ class ForwardBatch:
|
||||
# Image inputs
|
||||
image_path: Optional[str] = None
|
||||
image_embeds: List[torch.Tensor] = field(default_factory=list)
|
||||
pil_image: Optional[PIL.Image.Image] = None
|
||||
preprocessed_image: Optional[torch.Tensor] = None
|
||||
|
||||
# Text inputs
|
||||
prompt: Optional[Union[str, List[str]]] = None
|
||||
@@ -136,3 +140,37 @@ class ForwardBatch:
|
||||
|
||||
def __str__(self):
|
||||
return pprint.pformat(asdict(self), indent=2, width=120)
|
||||
|
||||
|
||||
@dataclass
|
||||
class TrainingBatch:
|
||||
current_timestep: int = 0
|
||||
current_vsa_sparsity: float = 0.0
|
||||
|
||||
# Dataloader batch outputs
|
||||
latents: Optional[torch.Tensor] = None
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None
|
||||
encoder_attention_mask: Optional[torch.Tensor] = None
|
||||
# i2v
|
||||
preprocessed_image: Optional[torch.Tensor] = None
|
||||
image_embeds: Optional[torch.Tensor] = None
|
||||
image_latents: Optional[torch.Tensor] = None
|
||||
infos: Optional[List[Dict[str, Any]]] = None
|
||||
|
||||
# Transformer inputs
|
||||
noisy_model_input: Optional[torch.Tensor] = None
|
||||
timesteps: Optional[torch.Tensor] = None
|
||||
sigmas: Optional[torch.Tensor] = None
|
||||
noise: Optional[torch.Tensor] = None
|
||||
|
||||
attn_metadata: Optional[AttentionMetadata] = None
|
||||
|
||||
# input kwargs
|
||||
input_kwargs: Optional[Dict[str, Any]] = None
|
||||
|
||||
# Training loss
|
||||
loss: torch.Tensor | None = None
|
||||
|
||||
# Training outputs
|
||||
total_loss: float | None = None
|
||||
grad_norm: float | None = None
|
||||
|
||||
@@ -2,19 +2,22 @@
|
||||
import gc
|
||||
import multiprocessing
|
||||
import os
|
||||
from collections import defaultdict
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
from itertools import chain
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.dataset import getdataset
|
||||
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.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -46,7 +49,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
# Initialize class variables for data sharing
|
||||
self.video_data: Dict[str, Any] = {} # Store video metadata and paths
|
||||
self.latent_data: Dict[str, Any] = {} # Store latent tensors
|
||||
self.preprocess_validation_text(fastvideo_args, args)
|
||||
self.preprocess_validation(fastvideo_args, args)
|
||||
self.preprocess_video_and_text(fastvideo_args, args)
|
||||
|
||||
def get_extra_features(self, valid_data: Dict[str, Any],
|
||||
@@ -58,39 +61,206 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
"""Get the schema fields for the pipeline type. Override in subclasses."""
|
||||
raise NotImplementedError
|
||||
|
||||
def create_record_for_schema(self,
|
||||
preprocess_batch: PreprocessBatch,
|
||||
schema: pa.Schema,
|
||||
strict: bool = False) -> Dict[str, Any]:
|
||||
"""Create a record for the Parquet dataset using a generic schema-based approach.
|
||||
|
||||
Args:
|
||||
preprocess_batch: The batch containing the data to extract
|
||||
schema: PyArrow schema defining the expected fields
|
||||
strict: If True, raises an exception when required fields are missing or unfilled
|
||||
|
||||
Returns:
|
||||
Dictionary record matching the schema
|
||||
|
||||
Raises:
|
||||
ValueError: If strict=True and required fields are missing or unfilled
|
||||
"""
|
||||
record = {}
|
||||
unfilled_fields = []
|
||||
|
||||
for field in schema.names:
|
||||
field_filled = False
|
||||
|
||||
if field.endswith('_bytes'):
|
||||
# Handle binary tensor data - convert numpy array or tensor to bytes
|
||||
tensor_name = field.replace('_bytes', '')
|
||||
tensor_data = getattr(preprocess_batch, tensor_name, None)
|
||||
if tensor_data is not None:
|
||||
try:
|
||||
if hasattr(tensor_data, 'numpy'): # torch tensor
|
||||
record[field] = tensor_data.cpu().numpy().tobytes()
|
||||
field_filled = True
|
||||
elif hasattr(tensor_data, 'tobytes'): # numpy array
|
||||
record[field] = tensor_data.tobytes()
|
||||
field_filled = True
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported tensor type for field {field}: {type(tensor_data)}"
|
||||
)
|
||||
except Exception as e:
|
||||
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
|
||||
|
||||
elif field.endswith('_shape'):
|
||||
# Handle tensor shape info
|
||||
tensor_name = field.replace('_shape', '')
|
||||
tensor_data = getattr(preprocess_batch, tensor_name, None)
|
||||
if tensor_data is not None and hasattr(tensor_data, 'shape'):
|
||||
record[field] = list(tensor_data.shape)
|
||||
field_filled = True
|
||||
else:
|
||||
record[field] = []
|
||||
|
||||
elif field.endswith('_dtype'):
|
||||
# Handle tensor dtype info
|
||||
tensor_name = field.replace('_dtype', '')
|
||||
tensor_data = getattr(preprocess_batch, tensor_name, None)
|
||||
if tensor_data is not None and hasattr(tensor_data, 'dtype'):
|
||||
record[field] = str(tensor_data.dtype)
|
||||
field_filled = True
|
||||
else:
|
||||
record[field] = 'unknown'
|
||||
|
||||
elif field in ['width', 'height', 'num_frames']:
|
||||
# Handle integer metadata fields
|
||||
value = getattr(preprocess_batch, field, None)
|
||||
if value is not None:
|
||||
try:
|
||||
record[field] = int(value)
|
||||
field_filled = True
|
||||
except (ValueError, TypeError) as e:
|
||||
if strict:
|
||||
raise ValueError(
|
||||
f"Failed to convert field {field} to int: {e}"
|
||||
) from e
|
||||
record[field] = 0
|
||||
else:
|
||||
record[field] = 0
|
||||
|
||||
elif field in ['duration_sec', 'fps']:
|
||||
# Handle float metadata fields
|
||||
# Map schema field names to batch attribute names
|
||||
attr_name = 'duration' if field == 'duration_sec' else field
|
||||
value = getattr(preprocess_batch, attr_name, None)
|
||||
if value is not None:
|
||||
try:
|
||||
record[field] = float(value)
|
||||
field_filled = True
|
||||
except (ValueError, TypeError) as e:
|
||||
if strict:
|
||||
raise ValueError(
|
||||
f"Failed to convert field {field} to float: {e}"
|
||||
) from e
|
||||
record[field] = 0.0
|
||||
else:
|
||||
record[field] = 0.0
|
||||
|
||||
else:
|
||||
# Handle string fields (id, file_name, caption, media_type, etc.)
|
||||
# Map common schema field names to batch attribute names
|
||||
attr_name = field
|
||||
if field == 'caption':
|
||||
attr_name = 'text'
|
||||
elif field == 'file_name':
|
||||
attr_name = 'path'
|
||||
elif field == 'id':
|
||||
# Generate ID from path if available
|
||||
path_value = getattr(preprocess_batch, 'path', None)
|
||||
if path_value:
|
||||
import os
|
||||
record[field] = os.path.basename(path_value).split(
|
||||
'.')[0]
|
||||
field_filled = True
|
||||
else:
|
||||
record[field] = ""
|
||||
continue
|
||||
elif field == 'media_type':
|
||||
# Determine media type from path
|
||||
path_value = getattr(preprocess_batch, 'path', None)
|
||||
if path_value:
|
||||
record[field] = 'video' if path_value.endswith(
|
||||
'.mp4') else 'image'
|
||||
field_filled = True
|
||||
else:
|
||||
record[field] = ""
|
||||
continue
|
||||
|
||||
value = getattr(preprocess_batch, attr_name, None)
|
||||
if value is not None:
|
||||
record[field] = str(value)
|
||||
field_filled = True
|
||||
else:
|
||||
record[field] = ""
|
||||
|
||||
# Track unfilled fields
|
||||
if not field_filled:
|
||||
unfilled_fields.append(field)
|
||||
|
||||
# Handle strict mode
|
||||
if strict and unfilled_fields:
|
||||
raise ValueError(
|
||||
f"Required fields were not filled: {unfilled_fields}")
|
||||
|
||||
# Log unfilled fields as warning if not in strict mode
|
||||
if unfilled_fields:
|
||||
logger.warning(
|
||||
"Some fields were not filled and got default values: %s",
|
||||
unfilled_fields)
|
||||
|
||||
return record
|
||||
|
||||
def create_record(
|
||||
self,
|
||||
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."""
|
||||
record = {
|
||||
"id": video_name,
|
||||
"vae_latent_bytes": vae_latent.tobytes(),
|
||||
"vae_latent_shape": list(vae_latent.shape),
|
||||
"vae_latent_dtype": str(vae_latent.dtype),
|
||||
"text_embedding_bytes": text_embedding.tobytes(),
|
||||
"text_embedding_shape": list(text_embedding.shape),
|
||||
"text_embedding_dtype": str(text_embedding.dtype),
|
||||
"text_attention_mask_bytes": text_attention_mask.tobytes(),
|
||||
"text_attention_mask_shape": list(text_attention_mask.shape),
|
||||
"text_attention_mask_dtype": str(text_attention_mask.dtype),
|
||||
"file_name": video_name,
|
||||
"caption": valid_data["text"][idx] if valid_data else "",
|
||||
"media_type": "video",
|
||||
"id":
|
||||
video_name,
|
||||
"vae_latent_bytes":
|
||||
vae_latent.tobytes(),
|
||||
"vae_latent_shape":
|
||||
list(vae_latent.shape),
|
||||
"vae_latent_dtype":
|
||||
str(vae_latent.dtype),
|
||||
"text_embedding_bytes":
|
||||
text_embedding.tobytes(),
|
||||
"text_embedding_shape":
|
||||
list(text_embedding.shape),
|
||||
"text_embedding_dtype":
|
||||
str(text_embedding.dtype),
|
||||
"file_name":
|
||||
video_name,
|
||||
"caption":
|
||||
valid_data["text"][idx] if len(valid_data["text"]) > 0 else "",
|
||||
"media_type":
|
||||
"video",
|
||||
"width":
|
||||
valid_data["pixel_values"][idx].shape[-2] if valid_data else 0,
|
||||
valid_data["pixel_values"][idx].shape[-2]
|
||||
if len(valid_data["pixel_values"]) > 0 else 0,
|
||||
"height":
|
||||
valid_data["pixel_values"][idx].shape[-1] if valid_data else 0,
|
||||
valid_data["pixel_values"][idx].shape[-1]
|
||||
if len(valid_data["pixel_values"]) > 0 else 0,
|
||||
"num_frames":
|
||||
vae_latent.shape[1] if len(vae_latent.shape) > 1 else 0,
|
||||
"duration_sec":
|
||||
float(valid_data["duration"][idx]) if valid_data else 0.0,
|
||||
"fps": float(valid_data["fps"][idx]) if valid_data else 0.0,
|
||||
float(valid_data["duration"][idx])
|
||||
if len(valid_data["duration"]) > 0 else 0.0,
|
||||
"fps":
|
||||
float(valid_data["fps"][idx])
|
||||
if len(valid_data["fps"]) > 0 else 0.0,
|
||||
}
|
||||
if extra_features:
|
||||
record.update(extra_features)
|
||||
@@ -103,7 +273,6 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
"combined_parquet_dataset")
|
||||
os.makedirs(combined_parquet_dir, exist_ok=True)
|
||||
local_rank = int(os.getenv("RANK", 0))
|
||||
world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
|
||||
# Get how many samples have already been processed
|
||||
start_idx = 0
|
||||
@@ -114,14 +283,10 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
start_idx += table.num_rows
|
||||
|
||||
# Loading dataset
|
||||
train_dataset = getdataset(args, start_idx=start_idx)
|
||||
sampler = DistributedSampler(train_dataset,
|
||||
rank=local_rank,
|
||||
num_replicas=world_size,
|
||||
shuffle=False)
|
||||
train_dataset = getdataset(args)
|
||||
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
batch_size=args.preprocess_video_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
@@ -215,8 +380,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 = {}
|
||||
@@ -233,7 +396,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)
|
||||
@@ -285,7 +447,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
num_processed_samples = 0
|
||||
self.all_tables = []
|
||||
|
||||
def preprocess_validation_text(self, fastvideo_args: FastVideoArgs, args):
|
||||
def preprocess_validation(self, fastvideo_args: FastVideoArgs, args):
|
||||
"""Process validation text prompts and save them to parquet files.
|
||||
|
||||
This base implementation handles the common validation text processing logic.
|
||||
@@ -296,22 +458,32 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
"validation_parquet_dataset")
|
||||
os.makedirs(validation_parquet_dir, exist_ok=True)
|
||||
|
||||
with open(args.validation_prompt_txt, encoding="utf-8") as file:
|
||||
lines = file.readlines()
|
||||
prompts = [line.strip() for line in lines]
|
||||
validation_dataset = ValidationDataset(args.validation_dataset_file)
|
||||
|
||||
# Prepare batch data for Parquet dataset
|
||||
batch_data = []
|
||||
sampling_param = SamplingParam.from_pretrained(
|
||||
fastvideo_args.model_path)
|
||||
if sampling_param.negative_prompt:
|
||||
prompts = [sampling_param.negative_prompt] + prompts
|
||||
negative_prompt = {
|
||||
'caption': sampling_param.negative_prompt,
|
||||
'image_path': None,
|
||||
'video_path': None,
|
||||
}
|
||||
validation_iterable = chain([negative_prompt], validation_dataset)
|
||||
else:
|
||||
negative_prompt = None
|
||||
validation_iterable = validation_dataset
|
||||
|
||||
# Add progress bar for validation text preprocessing
|
||||
pbar = tqdm(enumerate(prompts),
|
||||
pbar = tqdm(enumerate(validation_iterable),
|
||||
desc="Processing validation prompts",
|
||||
unit="prompt")
|
||||
for prompt_idx, prompt in pbar:
|
||||
for idx, sample in pbar:
|
||||
with torch.inference_mode():
|
||||
prompt = sample["caption"]
|
||||
is_negative_prompt = idx == 0
|
||||
|
||||
# Text Encoder
|
||||
batch = ForwardBatch(
|
||||
data_type="video",
|
||||
@@ -338,15 +510,43 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
"Shape after removing padding - Embeddings: %s, Mask: %s",
|
||||
text_embedding.shape, text_attention_mask.shape)
|
||||
|
||||
extra_features = {}
|
||||
if not is_negative_prompt:
|
||||
height = sample["height"]
|
||||
width = sample["width"]
|
||||
if "image_path" in sample and "video_path" in sample:
|
||||
raise ValueError(
|
||||
"Only one of image_path or video_path should be provided"
|
||||
)
|
||||
|
||||
if "image" in sample:
|
||||
extra_features = self.preprocess_image(
|
||||
sample["image"], height, width, fastvideo_args)
|
||||
|
||||
if "video" in sample:
|
||||
extra_features = self.preprocess_video(
|
||||
sample["video"], height, width, fastvideo_args)
|
||||
|
||||
# Get extra features for this sample if needed
|
||||
sample_extra_features = {}
|
||||
if extra_features:
|
||||
for key, value in extra_features.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
sample_extra_features[key] = value.cpu().numpy()
|
||||
else:
|
||||
sample_extra_features[key] = value
|
||||
|
||||
valid_data = defaultdict(list)
|
||||
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=None,
|
||||
valid_data=valid_data,
|
||||
idx=0,
|
||||
extra_features=None)
|
||||
extra_features=sample_extra_features)
|
||||
batch_data.append(record)
|
||||
|
||||
logger.info("Saved validation sample: %s", file_name)
|
||||
@@ -420,6 +620,15 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
del table
|
||||
gc.collect() # Force garbage collection
|
||||
|
||||
def preprocess_image(self, image: PIL.Image.Image, height: int, width: int,
|
||||
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
return {}
|
||||
|
||||
def preprocess_video(self, video: list[PIL.Image.Image], height: int,
|
||||
width: int,
|
||||
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
return {}
|
||||
|
||||
def _flush_tables(self, num_processed_samples: int, args,
|
||||
combined_parquet_dir: str):
|
||||
"""Flush collected tables to disk."""
|
||||
|
||||
@@ -8,6 +8,7 @@ using the modular pipeline architecture.
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import numpy as np
|
||||
import PIL
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
@@ -15,8 +16,13 @@ from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_i2v
|
||||
from fastvideo.v1.distributed import get_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,
|
||||
normalize, numpy_to_pt,
|
||||
pil_to_numpy, resize)
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_base import (
|
||||
BasePreprocessPipeline)
|
||||
from fastvideo.v1.pipelines.stages import ImageEncodingStage, TextEncodingStage
|
||||
|
||||
|
||||
class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
@@ -26,18 +32,73 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
"text_encoder", "tokenizer", "vae", "image_encoder", "image_processor"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="image_encoding_stage",
|
||||
stage=ImageEncodingStage(
|
||||
image_encoder=self.get_module("image_encoder"),
|
||||
image_processor=self.get_module("image_processor"),
|
||||
))
|
||||
|
||||
def preprocess_image(self, image: PIL.Image.Image, height: int, width: int,
|
||||
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
assert hasattr(
|
||||
self,
|
||||
"image_encoding_stage"), "Image encoding stage must be created"
|
||||
|
||||
batch = ForwardBatch(
|
||||
data_type="video",
|
||||
pil_image=image,
|
||||
)
|
||||
result_batch = self.image_encoding_stage(batch, fastvideo_args)
|
||||
clip_features = result_batch.image_embeds[0]
|
||||
|
||||
image = self.preprocess(
|
||||
image,
|
||||
vae_scale_factor=self.get_module("vae").spatial_compression_ratio,
|
||||
height=height,
|
||||
width=width)
|
||||
|
||||
return {
|
||||
"clip_feature": clip_features[0],
|
||||
"pil_image": image,
|
||||
}
|
||||
|
||||
def preprocess_video(self, video: list[PIL.Image.Image], height: int,
|
||||
width: int,
|
||||
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
return self.preprocess_image(video[0], height, width, fastvideo_args)
|
||||
|
||||
def get_schema_fields(self) -> List[str]:
|
||||
"""Get the schema fields for I2V pipeline."""
|
||||
return [f.name for f in pyarrow_schema_i2v]
|
||||
|
||||
def get_extra_features(self, valid_data: Dict[str, Any],
|
||||
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("vae").to(get_torch_device())
|
||||
|
||||
features = {}
|
||||
"""Get CLIP features from the first frame of each video."""
|
||||
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
|
||||
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
|
||||
batch_size, _, num_frames, height, width = valid_data[
|
||||
"pixel_values"].shape
|
||||
latent_height = height // self.get_module(
|
||||
"vae").spatial_compression_ratio
|
||||
latent_width = width // self.get_module("vae").spatial_compression_ratio
|
||||
|
||||
processed_images = []
|
||||
# Frame has values between -1 and 1
|
||||
for frame in first_frame:
|
||||
frame = (frame + 1) * 127.5
|
||||
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
|
||||
processed_img = self.get_module("image_processor")(
|
||||
images=frame_pil, return_tensors="pt")
|
||||
@@ -53,22 +114,84 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
clip_features = self.get_module("image_encoder")(**image_inputs)
|
||||
clip_features = clip_features.last_hidden_state
|
||||
|
||||
return {"clip_feature": clip_features}
|
||||
features["clip_feature"] = clip_features
|
||||
"""Get VAE features from the first frame of each video"""
|
||||
video_conditions = []
|
||||
for frame in first_frame:
|
||||
processed_img = frame.to(device="cpu", dtype=torch.float32)
|
||||
processed_img = processed_img.unsqueeze(0).permute(0, 3, 1,
|
||||
2).unsqueeze(2)
|
||||
# (B, H, W, C) -> (B, C, 1, H, W)
|
||||
video_condition = torch.cat([
|
||||
processed_img,
|
||||
processed_img.new_zeros(processed_img.shape[0],
|
||||
processed_img.shape[1], num_frames - 1,
|
||||
height, width)
|
||||
],
|
||||
dim=2)
|
||||
video_condition = video_condition.to(device=get_torch_device(),
|
||||
dtype=torch.float32)
|
||||
video_conditions.append(video_condition)
|
||||
|
||||
video_conditions = torch.cat(video_conditions, dim=0)
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=torch.float32,
|
||||
enabled=True):
|
||||
encoder_outputs = self.get_module("vae").encode(video_conditions)
|
||||
|
||||
latent_condition = encoder_outputs.mean
|
||||
if (hasattr(self.get_module("vae"), "shift_factor")
|
||||
and self.get_module("vae").shift_factor is not None):
|
||||
if isinstance(self.get_module("vae").shift_factor, torch.Tensor):
|
||||
latent_condition -= self.get_module("vae").shift_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition -= self.get_module("vae").shift_factor
|
||||
|
||||
if isinstance(self.get_module("vae").scaling_factor, torch.Tensor):
|
||||
latent_condition = latent_condition * self.get_module(
|
||||
"vae").scaling_factor.to(latent_condition.device,
|
||||
latent_condition.dtype)
|
||||
else:
|
||||
latent_condition = latent_condition * self.get_module(
|
||||
"vae").scaling_factor
|
||||
|
||||
mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
|
||||
latent_width)
|
||||
mask_lat_size[:, :, list(range(1, num_frames))] = 0
|
||||
first_frame_mask = mask_lat_size[:, :, 0:1]
|
||||
first_frame_mask = torch.repeat_interleave(
|
||||
first_frame_mask,
|
||||
dim=2,
|
||||
repeats=self.get_module("vae").temporal_compression_ratio)
|
||||
mask_lat_size = torch.concat(
|
||||
[first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
|
||||
mask_lat_size = mask_lat_size.view(
|
||||
batch_size, -1,
|
||||
self.get_module("vae").temporal_compression_ratio, latent_height,
|
||||
latent_width)
|
||||
mask_lat_size = mask_lat_size.transpose(1, 2)
|
||||
mask_lat_size = mask_lat_size.to(latent_condition.device)
|
||||
|
||||
image_latent = torch.concat([mask_lat_size, latent_condition], dim=1)
|
||||
|
||||
features["first_frame_latent"] = image_latent
|
||||
|
||||
return features
|
||||
|
||||
def create_record(
|
||||
self,
|
||||
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)
|
||||
@@ -87,7 +210,69 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
"clip_feature_dtype": "",
|
||||
})
|
||||
|
||||
return record # type: ignore
|
||||
if extra_features and "first_frame_latent" in extra_features:
|
||||
first_frame_latent = extra_features["first_frame_latent"]
|
||||
record.update({
|
||||
"first_frame_latent_bytes":
|
||||
first_frame_latent.tobytes(),
|
||||
"first_frame_latent_shape":
|
||||
list(first_frame_latent.shape),
|
||||
"first_frame_latent_dtype":
|
||||
str(first_frame_latent.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"first_frame_latent_bytes": b"",
|
||||
"first_frame_latent_shape": [],
|
||||
"first_frame_latent_dtype": "",
|
||||
})
|
||||
|
||||
if extra_features and "pil_image" in extra_features:
|
||||
pil_image = extra_features["pil_image"]
|
||||
record.update({
|
||||
"pil_image_bytes": pil_image.tobytes(),
|
||||
"pil_image_shape": list(pil_image.shape),
|
||||
"pil_image_dtype": str(pil_image.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"pil_image_bytes": b"",
|
||||
"pil_image_shape": [],
|
||||
"pil_image_dtype": "",
|
||||
})
|
||||
|
||||
return record
|
||||
|
||||
def pil_to_tensor(self, image: PIL.Image.Image) -> torch.Tensor:
|
||||
image = image
|
||||
|
||||
image = np.array(image).astype(np.float32)
|
||||
image = torch.from_numpy(image)
|
||||
return image
|
||||
|
||||
def preprocess(self,
|
||||
image: PIL.Image.Image,
|
||||
vae_scale_factor: int,
|
||||
height: int,
|
||||
width: int,
|
||||
resize_mode: str = "default") -> torch.Tensor:
|
||||
image = [image]
|
||||
|
||||
height, width = get_default_height_width(image[0], vae_scale_factor,
|
||||
height, width)
|
||||
image = [
|
||||
resize(i, height, width, resize_mode=resize_mode) for i in image
|
||||
]
|
||||
image = pil_to_numpy(image) # to np
|
||||
image = numpy_to_pt(image) # to pt
|
||||
|
||||
do_normalize = True
|
||||
if image.min() < 0:
|
||||
do_normalize = False
|
||||
if do_normalize:
|
||||
image = normalize(image)
|
||||
|
||||
return image
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_I2V
|
||||
|
||||
@@ -44,7 +44,7 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--model_type", type=str, default="mochi")
|
||||
parser.add_argument("--data_merge_path", type=str, required=True)
|
||||
parser.add_argument("--validation_prompt_txt", type=str)
|
||||
parser.add_argument("--validation_dataset_file", type=str)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument(
|
||||
"--dataloader_num_workers",
|
||||
@@ -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,
|
||||
@@ -79,7 +73,6 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--dataset", default="t2v")
|
||||
parser.add_argument("--preprocess_task", type=str, default="t2v")
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
@@ -91,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,
|
||||
|
||||
@@ -15,10 +15,16 @@ import torch
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class StageVerificationError(Exception):
|
||||
"""Exception raised when stage verification fails."""
|
||||
pass
|
||||
|
||||
|
||||
class PipelineStage(ABC):
|
||||
"""
|
||||
Abstract base class for all pipeline stages.
|
||||
@@ -28,6 +34,70 @@ class PipelineStage(ABC):
|
||||
for a specific part of the process, such as prompt encoding, latent preparation, etc.
|
||||
"""
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""
|
||||
Verify the input for the stage.
|
||||
|
||||
Example:
|
||||
from fastvideo.v1.pipelines.stages.validators import V, VerificationResult
|
||||
|
||||
def verify_input(self, batch, fastvideo_args):
|
||||
result = VerificationResult()
|
||||
result.add_check("height", batch.height, V.positive_int_divisible(8))
|
||||
result.add_check("width", batch.width, V.positive_int_divisible(8))
|
||||
result.add_check("image_latent", batch.image_latent, V.is_tensor)
|
||||
return result
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
A VerificationResult containing the verification status.
|
||||
|
||||
"""
|
||||
# Default implementation - no verification
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""
|
||||
Verify the output for the stage.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
A VerificationResult containing the verification status.
|
||||
"""
|
||||
# Default implementation - no verification
|
||||
return VerificationResult()
|
||||
|
||||
def _run_verification(self, verification_result: VerificationResult,
|
||||
stage_name: str, verification_type: str) -> None:
|
||||
"""
|
||||
Run verification and raise errors if any checks fail.
|
||||
|
||||
Args:
|
||||
verification_result: Results from verify_input or verify_output
|
||||
stage_name: Name of the current stage
|
||||
verification_type: "input" or "output"
|
||||
"""
|
||||
if not verification_result.is_valid():
|
||||
failed_fields = verification_result.get_failed_fields()
|
||||
if failed_fields:
|
||||
# Get detailed failure information
|
||||
detailed_summary = verification_result.get_failure_summary()
|
||||
|
||||
failed_fields_str = ", ".join(failed_fields)
|
||||
error_msg = (
|
||||
f"{verification_type.capitalize()} verification failed for {stage_name}: "
|
||||
f"Failed fields: {failed_fields_str}\n"
|
||||
f"Details: {detailed_summary}")
|
||||
raise StageVerificationError(error_msg)
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
"""Get the device for this stage."""
|
||||
@@ -48,7 +118,7 @@ class PipelineStage(ABC):
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Execute the stage's processing on the batch with optional logging.
|
||||
Execute the stage's processing on the batch with optional verification and logging.
|
||||
Should not be overridden by subclasses.
|
||||
|
||||
Args:
|
||||
@@ -58,34 +128,56 @@ class PipelineStage(ABC):
|
||||
Returns:
|
||||
The updated batch information after this stage's processing.
|
||||
"""
|
||||
# if envs.ENABLE_STAGE_LOGGING:
|
||||
stage_name = self.__class__.__name__
|
||||
|
||||
# Check if verification is enabled (simple approach for prototype)
|
||||
enable_verification = getattr(fastvideo_args,
|
||||
'enable_stage_verification', False)
|
||||
|
||||
if enable_verification:
|
||||
# Pre-execution input verification
|
||||
try:
|
||||
input_result = self.verify_input(batch, fastvideo_args)
|
||||
self._run_verification(input_result, stage_name, "input")
|
||||
except Exception as e:
|
||||
logger.error("Input verification failed for %s: %s", stage_name,
|
||||
str(e))
|
||||
raise
|
||||
|
||||
# Execute the actual stage logic
|
||||
# envs.ENABLE_STAGE_LOGGING
|
||||
if False:
|
||||
self._logger.info("[%s] Starting execution", self._stage_name)
|
||||
self._logger.info("[%s] Starting execution", stage_name)
|
||||
start_time = time.perf_counter()
|
||||
|
||||
try:
|
||||
# Call the actual implementation
|
||||
result = self._call_implementation(batch, fastvideo_args)
|
||||
|
||||
result = self.forward(batch, fastvideo_args)
|
||||
execution_time = time.perf_counter() - start_time
|
||||
self._logger.info("[%s] Execution completed in %s ms",
|
||||
self._stage_name, execution_time * 1000)
|
||||
|
||||
return result
|
||||
stage_name, execution_time * 1000)
|
||||
except Exception as e:
|
||||
execution_time = time.perf_counter() - start_time
|
||||
self._logger.error(
|
||||
"[%s] Error during execution after %s ms: %s",
|
||||
self._stage_name, execution_time * 1000, e)
|
||||
self._logger.error("[%s] Traceback: %s", self._stage_name,
|
||||
"[%s] Error during execution after %s ms: %s", stage_name,
|
||||
execution_time * 1000, e)
|
||||
self._logger.error("[%s] Traceback: %s", stage_name,
|
||||
traceback.format_exc())
|
||||
|
||||
# Re-raise the exception
|
||||
raise
|
||||
else:
|
||||
# Just call the implementation directly if logging is disabled
|
||||
# TODO(will): Also handle backward
|
||||
return self.forward(batch, fastvideo_args)
|
||||
# Direct execution (current behavior)
|
||||
result = self.forward(batch, fastvideo_args)
|
||||
|
||||
if enable_verification:
|
||||
# Post-execution output verification
|
||||
try:
|
||||
output_result = self.verify_output(result, fastvideo_args)
|
||||
self._run_verification(output_result, stage_name, "output")
|
||||
except Exception as e:
|
||||
logger.error("Output verification failed for %s: %s",
|
||||
stage_name, str(e))
|
||||
raise
|
||||
|
||||
return result
|
||||
|
||||
@abstractmethod
|
||||
def forward(
|
||||
|
||||
@@ -9,6 +9,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -69,3 +71,24 @@ class ConditioningStage(PipelineStage):
|
||||
[batch.negative_attention_mask_2, batch.attention_mask_2])
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify conditioning stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("do_classifier_free_guidance",
|
||||
batch.do_classifier_free_guidance, V.bool_value)
|
||||
result.add_check("guidance_scale", batch.guidance_scale,
|
||||
V.positive_float)
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
result.add_check(
|
||||
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
|
||||
not batch.do_classifier_free_guidance or V.list_not_empty(x))
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify conditioning stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
return result
|
||||
|
||||
@@ -11,6 +11,8 @@ from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -27,6 +29,23 @@ class DecodingStage(PipelineStage):
|
||||
def __init__(self, vae) -> None:
|
||||
self.vae: ParallelTiledVAE = vae
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify decoding stage inputs."""
|
||||
result = VerificationResult()
|
||||
# Denoised latents for VAE decoding: [batch_size, channels, frames, height_latents, width_latents]
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify decoding stage outputs."""
|
||||
result = VerificationResult()
|
||||
# Decoded video/images: [batch_size, channels, frames, height, width]
|
||||
result.add_check("output", batch.output, [V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
|
||||
@@ -3,15 +3,15 @@
|
||||
Denoising stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
import inspect
|
||||
from typing import Any, Dict, Iterable, List, Optional
|
||||
from typing import Any, Dict, Iterable, Optional
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
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.communication_op import (
|
||||
@@ -21,19 +21,24 @@ from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
from fastvideo.v1.utils import dict_to_3d_list
|
||||
|
||||
st_attn_available = False
|
||||
if importlib.util.find_spec("st_attn") is not None:
|
||||
st_attn_available = True
|
||||
try:
|
||||
from fastvideo.v1.attention.backends.sliding_tile_attn import (
|
||||
SlidingTileAttentionBackend)
|
||||
st_attn_available = True
|
||||
except ImportError:
|
||||
st_attn_available = False
|
||||
|
||||
vsa_available = False
|
||||
if importlib.util.find_spec("vsa") is not None:
|
||||
vsa_available = True
|
||||
try:
|
||||
from fastvideo.v1.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionBackend)
|
||||
vsa_available = True
|
||||
except ImportError:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -117,20 +122,6 @@ class DenoisingStage(PipelineStage):
|
||||
num_warmup_steps = len(
|
||||
timesteps) - num_inference_steps * self.scheduler.order
|
||||
|
||||
# Create 3D list for mask strategy
|
||||
def dict_to_3d_list(mask_strategy,
|
||||
t_max=50,
|
||||
l_max=60,
|
||||
h_max=24) -> List:
|
||||
result = [[[None for _ in range(h_max)] for _ in range(l_max)]
|
||||
for _ in range(t_max)]
|
||||
if mask_strategy is None:
|
||||
return result
|
||||
for key, value in mask_strategy.items():
|
||||
t, layer, h = map(int, key.split('_'))
|
||||
result[t][layer][h] = value
|
||||
return result
|
||||
|
||||
# Prepare image latents and embeddings for I2V generation
|
||||
image_embeds = batch.image_embeds
|
||||
if len(image_embeds) > 0:
|
||||
@@ -143,7 +134,8 @@ class DenoisingStage(PipelineStage):
|
||||
self.transformer.forward,
|
||||
{
|
||||
"encoder_hidden_states_image": image_embeds,
|
||||
"mask_strategy": dict_to_3d_list(None)
|
||||
"mask_strategy": dict_to_3d_list(
|
||||
None, t_max=50, l_max=60, h_max=24)
|
||||
},
|
||||
)
|
||||
|
||||
@@ -304,7 +296,7 @@ class DenoisingStage(PipelineStage):
|
||||
batch.latents = latents
|
||||
|
||||
# Save STA mask search results if needed
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == 'STA_searching':
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == STA_Mode.STA_SEARCHING:
|
||||
self.save_sta_search_results(batch)
|
||||
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
@@ -402,7 +394,7 @@ class DenoisingStage(PipelineStage):
|
||||
raise NotImplementedError(
|
||||
"STA mask search/tuning is not supported for this resolution")
|
||||
|
||||
if STA_mode == "STA_searching" or STA_mode == "STA_tuning" or STA_mode == "STA_tuning_cfg":
|
||||
if STA_mode == STA_Mode.STA_SEARCHING or STA_mode == STA_Mode.STA_TUNING or STA_mode == STA_Mode.STA_TUNING_CFG:
|
||||
size = (batch.width, batch.height)
|
||||
if size == (1280, 768):
|
||||
# TODO: make it configurable
|
||||
@@ -424,18 +416,18 @@ class DenoisingStage(PipelineStage):
|
||||
layer_num += self.transformer.config.num_single_layers
|
||||
head_num = self.transformer.config.num_attention_heads
|
||||
|
||||
if STA_mode == "STA_searching":
|
||||
if STA_mode == STA_Mode.STA_SEARCHING:
|
||||
STA_param = configure_sta(
|
||||
mode='STA_searching',
|
||||
mode=STA_Mode.STA_SEARCHING,
|
||||
layer_num=layer_num,
|
||||
head_num=head_num,
|
||||
time_step_num=timesteps_num,
|
||||
mask_candidates=sparse_mask_candidates_searching +
|
||||
full_mask, # last is full mask; Can add more sparse masks while keep last one as full mask
|
||||
)
|
||||
elif STA_mode == 'STA_tuning':
|
||||
elif STA_mode == STA_Mode.STA_TUNING:
|
||||
STA_param = configure_sta(
|
||||
mode='STA_tuning',
|
||||
mode=STA_Mode.STA_TUNING,
|
||||
layer_num=layer_num,
|
||||
head_num=head_num,
|
||||
time_step_num=timesteps_num,
|
||||
@@ -448,9 +440,9 @@ class DenoisingStage(PipelineStage):
|
||||
save_dir=
|
||||
f'output/mask_search_strategy_{size[0]}x{size[1]}/', # Custom save directory
|
||||
timesteps=timesteps_num)
|
||||
elif STA_mode == 'STA_tuning_cfg':
|
||||
elif STA_mode == STA_Mode.STA_TUNING_CFG:
|
||||
STA_param = configure_sta(
|
||||
mode='STA_tuning_cfg',
|
||||
mode=STA_Mode.STA_TUNING_CFG,
|
||||
layer_num=layer_num,
|
||||
head_num=head_num,
|
||||
time_step_num=timesteps_num,
|
||||
@@ -463,12 +455,12 @@ class DenoisingStage(PipelineStage):
|
||||
skip_time_steps=skip_time_steps,
|
||||
save_dir=f'output/mask_search_strategy_{size[0]}x{size[1]}/',
|
||||
timesteps=timesteps_num)
|
||||
elif STA_mode == 'STA_inference':
|
||||
elif STA_mode == STA_Mode.STA_INFERENCE:
|
||||
import fastvideo.v1.envs as envs
|
||||
config_file = envs.FASTVIDEO_ATTENTION_CONFIG
|
||||
if config_file is None:
|
||||
raise ValueError("FASTVIDEO_ATTENTION_CONFIG is not set")
|
||||
STA_param = configure_sta(mode='STA_inference',
|
||||
STA_param = configure_sta(mode=STA_Mode.STA_INFERENCE,
|
||||
layer_num=layer_num,
|
||||
head_num=head_num,
|
||||
time_step_num=timesteps_num,
|
||||
@@ -517,3 +509,37 @@ class DenoisingStage(PipelineStage):
|
||||
mask_strategies=sparse_mask_candidates_searching,
|
||||
output_dir=f'output/mask_search_result_neg_{size[0]}x{size[1]}/'
|
||||
)
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify denoising stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("timesteps", batch.timesteps,
|
||||
[V.is_tensor, V.min_dims(1)])
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
result.add_check("image_embeds", batch.image_embeds, V.is_list)
|
||||
result.add_check("image_latent", batch.image_latent,
|
||||
V.none_or_tensor_with_dims(5))
|
||||
result.add_check("num_inference_steps", batch.num_inference_steps,
|
||||
V.positive_int)
|
||||
result.add_check("guidance_scale", batch.guidance_scale,
|
||||
V.positive_float)
|
||||
result.add_check("eta", batch.eta, V.non_negative_float)
|
||||
result.add_check("generator", batch.generator,
|
||||
V.generator_or_list_generators)
|
||||
result.add_check("do_classifier_free_guidance",
|
||||
batch.do_classifier_free_guidance, V.bool_value)
|
||||
result.add_check(
|
||||
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
|
||||
not batch.do_classifier_free_guidance or V.list_not_empty(x))
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify denoising stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
@@ -12,10 +12,12 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.models.vision_utils import (get_default_height_width,
|
||||
load_image, normalize,
|
||||
numpy_to_pt, pil_to_numpy, resize)
|
||||
normalize, numpy_to_pt,
|
||||
pil_to_numpy, resize)
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import V # Import validators
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -49,22 +51,27 @@ class EncodingStage(PipelineStage):
|
||||
"""
|
||||
self.vae = self.vae.to(get_torch_device())
|
||||
|
||||
image_path = batch.image_path
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
if image_path is None:
|
||||
raise ValueError("Image Path must be provided")
|
||||
assert batch.height is not None
|
||||
assert batch.width is not None
|
||||
latent_height = batch.height // self.vae.spatial_compression_ratio
|
||||
latent_width = batch.width // self.vae.spatial_compression_ratio
|
||||
|
||||
image = load_image(image_path)
|
||||
image = self.preprocess(
|
||||
image,
|
||||
vae_scale_factor=self.vae.spatial_compression_ratio,
|
||||
height=batch.height,
|
||||
width=batch.width).to(get_torch_device(), dtype=torch.float32)
|
||||
image = image.unsqueeze(2)
|
||||
image = batch.preprocessed_image
|
||||
# TODO(will)
|
||||
if image is None:
|
||||
assert batch.pil_image is not None
|
||||
image = batch.pil_image
|
||||
image = self.preprocess(
|
||||
image,
|
||||
vae_scale_factor=self.vae.spatial_compression_ratio,
|
||||
height=batch.height,
|
||||
width=batch.width).to(get_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([
|
||||
image,
|
||||
image.new_zeros(image.shape[0], image.shape[1],
|
||||
@@ -95,7 +102,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[0])
|
||||
latent_condition = self.retrieve_latents(encoder_output, generator)
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
@@ -174,3 +181,23 @@ class EncodingStage(PipelineStage):
|
||||
image = normalize(image)
|
||||
|
||||
return image
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify encoding stage inputs."""
|
||||
result = VerificationResult()
|
||||
# result.add_check("pil_image", batch.pil_image)
|
||||
result.add_check("height", batch.height, V.positive_int)
|
||||
result.add_check("width", batch.width, V.positive_int)
|
||||
result.add_check("generator", batch.generator,
|
||||
V.generator_or_list_generators)
|
||||
result.add_check("num_frames", batch.num_frames, V.positive_int)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify encoding stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("image_latent", batch.image_latent,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
@@ -11,9 +11,10 @@ from fastvideo.v1.distributed import get_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
|
||||
from fastvideo.v1.models.vision_utils import load_image
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -56,7 +57,7 @@ class ImageEncodingStage(PipelineStage):
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.image_encoder = self.image_encoder.to(get_torch_device())
|
||||
|
||||
image = load_image(batch.image_path)
|
||||
image = batch.pil_image
|
||||
|
||||
image_inputs = self.image_processor(
|
||||
images=image, return_tensors="pt").to(get_torch_device())
|
||||
@@ -71,3 +72,19 @@ class ImageEncodingStage(PipelineStage):
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify image encoding stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("pil_image", batch.pil_image, V.not_none)
|
||||
result.add_check("image_embeds", batch.image_embeds, V.is_list)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify image encoding stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("image_embeds", batch.image_embeds,
|
||||
V.list_of_tensors_dims(3))
|
||||
return result
|
||||
|
||||
@@ -7,11 +7,17 @@ import torch
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vision_utils import load_image
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import (StageValidators,
|
||||
VerificationResult)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Alias for convenience
|
||||
V = StageValidators
|
||||
|
||||
|
||||
class InputValidationStage(PipelineStage):
|
||||
"""
|
||||
@@ -86,4 +92,37 @@ class InputValidationStage(PipelineStage):
|
||||
f"Guidance scale must be positive, but got {batch.guidance_scale}"
|
||||
)
|
||||
|
||||
# for i2v, get image from image_path
|
||||
if batch.image_path is not None:
|
||||
image = load_image(batch.image_path)
|
||||
batch.pil_image = image
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify input validation stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("seed", batch.seed, [V.not_none, V.positive_int])
|
||||
result.add_check("num_videos_per_prompt", batch.num_videos_per_prompt,
|
||||
V.positive_int)
|
||||
result.add_check(
|
||||
"prompt_or_embeds", None, lambda _: V.string_or_list_strings(
|
||||
batch.prompt) or V.list_not_empty(batch.prompt_embeds))
|
||||
result.add_check("height", batch.height, V.positive_int)
|
||||
result.add_check("width", batch.width, V.positive_int)
|
||||
result.add_check("num_inference_steps", batch.num_inference_steps,
|
||||
V.positive_int)
|
||||
result.add_check(
|
||||
"guidance_scale", batch.guidance_scale, lambda x: not batch.
|
||||
do_classifier_free_guidance or V.positive_float(x))
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify input validation stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("seeds", batch.seeds, V.list_not_empty)
|
||||
result.add_check("generator", batch.generator,
|
||||
V.generator_or_list_generators)
|
||||
return result
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
"""
|
||||
Latent preparation stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
@@ -9,6 +10,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -126,4 +129,32 @@ class LatentPreparationStage(PipelineStage):
|
||||
latent_num_frames = (video_length - 1) // temporal_scale_factor + 1
|
||||
else: # stepvideo only
|
||||
latent_num_frames = video_length // 17 * 3
|
||||
return latent_num_frames
|
||||
return int(latent_num_frames)
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify latent preparation stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check(
|
||||
"prompt_or_embeds", None, lambda _: V.string_or_list_strings(
|
||||
batch.prompt) or V.list_not_empty(batch.prompt_embeds))
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds,
|
||||
V.list_of_tensors)
|
||||
result.add_check("num_videos_per_prompt", batch.num_videos_per_prompt,
|
||||
V.positive_int)
|
||||
result.add_check("generator", batch.generator,
|
||||
V.generator_or_list_generators)
|
||||
result.add_check("num_frames", batch.num_frames, V.positive_int)
|
||||
result.add_check("height", batch.height, V.positive_int)
|
||||
result.add_check("width", batch.width, V.positive_int)
|
||||
result.add_check("latents", batch.latents, V.none_or_tensor)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify latent preparation stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("raw_latent_shape", batch.raw_latent_shape, V.is_tuple)
|
||||
return result
|
||||
|
||||
@@ -1,10 +1,14 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -47,3 +51,29 @@ class StepvideoPromptEncodingStage(PipelineStage):
|
||||
batch.clip_embedding_pos = pos_clip
|
||||
batch.clip_embedding_neg = neg_clip
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify stepvideo encoding stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("prompt", batch.prompt, V.string_not_empty)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify stepvideo encoding stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds,
|
||||
[V.is_tensor, V.with_dims(3)])
|
||||
result.add_check("negative_prompt_embeds", batch.negative_prompt_embeds,
|
||||
[V.is_tensor, V.with_dims(3)])
|
||||
result.add_check("prompt_attention_mask", batch.prompt_attention_mask,
|
||||
[V.is_tensor, V.with_dims(2)])
|
||||
result.add_check("negative_attention_mask",
|
||||
batch.negative_attention_mask,
|
||||
[V.is_tensor, V.with_dims(2)])
|
||||
result.add_check("clip_embedding_pos", batch.clip_embedding_pos,
|
||||
[V.is_tensor, V.with_dims(2)])
|
||||
result.add_check("clip_embedding_neg", batch.clip_embedding_neg,
|
||||
[V.is_tensor, V.with_dims(2)])
|
||||
return result
|
||||
|
||||
@@ -12,6 +12,8 @@ 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
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = (__name__)
|
||||
|
||||
@@ -113,3 +115,30 @@ class TextEncodingStage(PipelineStage):
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify text encoding stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("prompt", batch.prompt, V.string_or_list_strings)
|
||||
result.add_check(
|
||||
"negative_prompt", batch.negative_prompt, lambda x: not batch.
|
||||
do_classifier_free_guidance or V.string_not_empty(x))
|
||||
result.add_check("do_classifier_free_guidance",
|
||||
batch.do_classifier_free_guidance, V.bool_value)
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.is_list)
|
||||
result.add_check("negative_prompt_embeds", batch.negative_prompt_embeds,
|
||||
V.none_or_list)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify text encoding stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds,
|
||||
V.list_of_tensors_min_dims(2))
|
||||
result.add_check(
|
||||
"negative_prompt_embeds", batch.negative_prompt_embeds,
|
||||
lambda x: not batch.do_classifier_free_guidance or V.
|
||||
list_of_tensors_with_min_dims(x, 2))
|
||||
return result
|
||||
|
||||
@@ -12,6 +12,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -95,3 +97,22 @@ class TimestepPreparationStage(PipelineStage):
|
||||
batch.timesteps = timesteps
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify timestep preparation stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("num_inference_steps", batch.num_inference_steps,
|
||||
V.positive_int)
|
||||
result.add_check("timesteps", batch.timesteps, V.none_or_tensor)
|
||||
result.add_check("sigmas", batch.sigmas, V.none_or_list)
|
||||
result.add_check("n_tokens", batch.n_tokens, V.none_or_positive_int)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify timestep preparation stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("timesteps", batch.timesteps,
|
||||
[V.is_tensor, V.with_dims(1)])
|
||||
return result
|
||||
|
||||
@@ -0,0 +1,486 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Common validators for pipeline stage verification.
|
||||
|
||||
This module provides reusable validation functions that can be used across
|
||||
all pipeline stages for input/output verification.
|
||||
"""
|
||||
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class StageValidators:
|
||||
"""Common validators for pipeline stages."""
|
||||
|
||||
@staticmethod
|
||||
def not_none(value: Any) -> bool:
|
||||
"""Check if value is not None."""
|
||||
return value is not None
|
||||
|
||||
@staticmethod
|
||||
def positive_int(value: Any) -> bool:
|
||||
"""Check if value is a positive integer."""
|
||||
return isinstance(value, int) and value > 0
|
||||
|
||||
@staticmethod
|
||||
def positive_float(value: Any) -> bool:
|
||||
"""Check if value is a positive float."""
|
||||
return isinstance(value, (int, float)) and value > 0
|
||||
|
||||
@staticmethod
|
||||
def non_negative_float(value: Any) -> bool:
|
||||
"""Check if value is a non-negative float."""
|
||||
return isinstance(value, (int, float)) and value >= 0
|
||||
|
||||
@staticmethod
|
||||
def divisible_by(value: Any, divisor: int) -> bool:
|
||||
"""Check if value is divisible by divisor."""
|
||||
return value is not None and isinstance(value,
|
||||
int) and value % divisor == 0
|
||||
|
||||
@staticmethod
|
||||
def is_tensor(value: Any) -> bool:
|
||||
"""Check if value is a torch tensor and doesn't contain NaN values."""
|
||||
if not isinstance(value, torch.Tensor):
|
||||
return False
|
||||
return not torch.isnan(value).any().item()
|
||||
|
||||
@staticmethod
|
||||
def tensor_with_dims(value: Any, dims: int) -> bool:
|
||||
"""Check if value is a tensor with specific dimensions and no NaN values."""
|
||||
if not isinstance(value, torch.Tensor):
|
||||
return False
|
||||
if value.dim() != dims:
|
||||
return False
|
||||
return not torch.isnan(value).any().item()
|
||||
|
||||
@staticmethod
|
||||
def tensor_min_dims(value: Any, min_dims: int) -> bool:
|
||||
"""Check if value is a tensor with at least min_dims dimensions and no NaN values."""
|
||||
if not isinstance(value, torch.Tensor):
|
||||
return False
|
||||
if value.dim() < min_dims:
|
||||
return False
|
||||
return not torch.isnan(value).any().item()
|
||||
|
||||
@staticmethod
|
||||
def tensor_shape_matches(value: Any, expected_shape: tuple) -> bool:
|
||||
"""Check if tensor shape matches expected shape (None for any size) and no NaN values."""
|
||||
if not isinstance(value, torch.Tensor):
|
||||
return False
|
||||
if len(value.shape) != len(expected_shape):
|
||||
return False
|
||||
for actual, expected in zip(value.shape, expected_shape):
|
||||
if expected is not None and actual != expected:
|
||||
return False
|
||||
return not torch.isnan(value).any().item()
|
||||
|
||||
@staticmethod
|
||||
def list_not_empty(value: Any) -> bool:
|
||||
"""Check if value is a non-empty list."""
|
||||
return isinstance(value, list) and len(value) > 0
|
||||
|
||||
@staticmethod
|
||||
def list_length(value: Any, length: int) -> bool:
|
||||
"""Check if list has specific length."""
|
||||
return isinstance(value, list) and len(value) == length
|
||||
|
||||
@staticmethod
|
||||
def list_min_length(value: Any, min_length: int) -> bool:
|
||||
"""Check if list has at least min_length items."""
|
||||
return isinstance(value, list) and len(value) >= min_length
|
||||
|
||||
@staticmethod
|
||||
def string_not_empty(value: Any) -> bool:
|
||||
"""Check if value is a non-empty string."""
|
||||
return isinstance(value, str) and len(value.strip()) > 0
|
||||
|
||||
@staticmethod
|
||||
def string_or_list_strings(value: Any) -> bool:
|
||||
"""Check if value is a string or list of strings."""
|
||||
if isinstance(value, str):
|
||||
return True
|
||||
if isinstance(value, list):
|
||||
return all(isinstance(item, str) for item in value)
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def bool_value(value: Any) -> bool:
|
||||
"""Check if value is a boolean."""
|
||||
return isinstance(value, bool)
|
||||
|
||||
@staticmethod
|
||||
def generator_or_list_generators(value: Any) -> bool:
|
||||
"""Check if value is a Generator or list of Generators."""
|
||||
if isinstance(value, torch.Generator):
|
||||
return True
|
||||
if isinstance(value, list):
|
||||
return all(isinstance(item, torch.Generator) for item in value)
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def is_list(value: Any) -> bool:
|
||||
"""Check if value is a list (can be empty)."""
|
||||
return isinstance(value, list)
|
||||
|
||||
@staticmethod
|
||||
def is_tuple(value: Any) -> bool:
|
||||
"""Check if value is a tuple."""
|
||||
return isinstance(value, tuple)
|
||||
|
||||
@staticmethod
|
||||
def none_or_tensor(value: Any) -> bool:
|
||||
"""Check if value is None or a tensor without NaN values."""
|
||||
if value is None:
|
||||
return True
|
||||
if not isinstance(value, torch.Tensor):
|
||||
return False
|
||||
return not torch.isnan(value).any().item()
|
||||
|
||||
@staticmethod
|
||||
def list_of_tensors_with_dims(value: Any, dims: int) -> bool:
|
||||
"""Check if value is a non-empty list where all items are tensors with specific dimensions and no NaN values."""
|
||||
if not isinstance(value, list) or len(value) == 0:
|
||||
return False
|
||||
for item in value:
|
||||
if not isinstance(item, torch.Tensor):
|
||||
return False
|
||||
if item.dim() != dims:
|
||||
return False
|
||||
if torch.isnan(item).any().item():
|
||||
return False
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def list_of_tensors(value: Any) -> bool:
|
||||
"""Check if value is a non-empty list where all items are tensors without NaN values."""
|
||||
if not isinstance(value, list) or len(value) == 0:
|
||||
return False
|
||||
for item in value:
|
||||
if not isinstance(item, torch.Tensor):
|
||||
return False
|
||||
if torch.isnan(item).any().item():
|
||||
return False
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def list_of_tensors_with_min_dims(value: Any, min_dims: int) -> bool:
|
||||
"""Check if value is a non-empty list where all items are tensors with at least min_dims dimensions and no NaN values."""
|
||||
if not isinstance(value, list) or len(value) == 0:
|
||||
return False
|
||||
for item in value:
|
||||
if not isinstance(item, torch.Tensor):
|
||||
return False
|
||||
if item.dim() < min_dims:
|
||||
return False
|
||||
if torch.isnan(item).any().item():
|
||||
return False
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def none_or_tensor_with_dims(dims: int) -> Callable[[Any], bool]:
|
||||
"""Return a validator that checks if value is None or a tensor with specific dimensions and no NaN values."""
|
||||
|
||||
def validator(value: Any) -> bool:
|
||||
if value is None:
|
||||
return True
|
||||
if not isinstance(value, torch.Tensor):
|
||||
return False
|
||||
if value.dim() != dims:
|
||||
return False
|
||||
return not torch.isnan(value).any().item()
|
||||
|
||||
return validator
|
||||
|
||||
@staticmethod
|
||||
def none_or_list(value: Any) -> bool:
|
||||
"""Check if value is None or a list."""
|
||||
return value is None or isinstance(value, list)
|
||||
|
||||
@staticmethod
|
||||
def none_or_positive_int(value: Any) -> bool:
|
||||
"""Check if value is None or a positive integer."""
|
||||
return value is None or (isinstance(value, int) and value > 0)
|
||||
|
||||
# Helper methods that return functions for common patterns
|
||||
@staticmethod
|
||||
def with_dims(dims: int) -> Callable[[Any], bool]:
|
||||
"""Return a validator that checks if tensor has specific dimensions and no NaN values."""
|
||||
|
||||
def validator(value: Any) -> bool:
|
||||
return StageValidators.tensor_with_dims(value, dims)
|
||||
|
||||
return validator
|
||||
|
||||
@staticmethod
|
||||
def min_dims(min_dims: int) -> Callable[[Any], bool]:
|
||||
"""Return a validator that checks if tensor has at least min_dims dimensions and no NaN values."""
|
||||
|
||||
def validator(value: Any) -> bool:
|
||||
return StageValidators.tensor_min_dims(value, min_dims)
|
||||
|
||||
return validator
|
||||
|
||||
@staticmethod
|
||||
def divisible(divisor: int) -> Callable[[Any], bool]:
|
||||
"""Return a validator that checks if value is divisible by divisor."""
|
||||
|
||||
def validator(value: Any) -> bool:
|
||||
return StageValidators.divisible_by(value, divisor)
|
||||
|
||||
return validator
|
||||
|
||||
@staticmethod
|
||||
def positive_int_divisible(divisor: int) -> Callable[[Any], bool]:
|
||||
"""Return a validator that checks if value is a positive integer divisible by divisor."""
|
||||
|
||||
def validator(value: Any) -> bool:
|
||||
return (isinstance(value, int) and value > 0
|
||||
and StageValidators.divisible_by(value, divisor))
|
||||
|
||||
return validator
|
||||
|
||||
@staticmethod
|
||||
def list_of_tensors_dims(dims: int) -> Callable[[Any], bool]:
|
||||
"""Return a validator that checks if value is a list of tensors with specific dimensions and no NaN values."""
|
||||
|
||||
def validator(value: Any) -> bool:
|
||||
return StageValidators.list_of_tensors_with_dims(value, dims)
|
||||
|
||||
return validator
|
||||
|
||||
@staticmethod
|
||||
def list_of_tensors_min_dims(min_dims: int) -> Callable[[Any], bool]:
|
||||
"""Return a validator that checks if value is a list of tensors with at least min_dims dimensions and no NaN values."""
|
||||
|
||||
def validator(value: Any) -> bool:
|
||||
return StageValidators.list_of_tensors_with_min_dims(
|
||||
value, min_dims)
|
||||
|
||||
return validator
|
||||
|
||||
|
||||
class ValidationFailure:
|
||||
"""Details about a specific validation failure."""
|
||||
|
||||
def __init__(self,
|
||||
validator_name: str,
|
||||
actual_value: Any,
|
||||
expected: Optional[str] = None,
|
||||
error_msg: Optional[str] = None):
|
||||
self.validator_name = validator_name
|
||||
self.actual_value = actual_value
|
||||
self.expected = expected
|
||||
self.error_msg = error_msg
|
||||
|
||||
def __str__(self) -> str:
|
||||
parts = [f"Validator '{self.validator_name}' failed"]
|
||||
|
||||
if self.error_msg:
|
||||
parts.append(f"Error: {self.error_msg}")
|
||||
|
||||
# Add actual value info (but limit very long representations)
|
||||
actual_str = self._format_value(self.actual_value)
|
||||
parts.append(f"Actual: {actual_str}")
|
||||
|
||||
if self.expected:
|
||||
parts.append(f"Expected: {self.expected}")
|
||||
|
||||
return ". ".join(parts)
|
||||
|
||||
def _format_value(self, value: Any) -> str:
|
||||
"""Format a value for display in error messages."""
|
||||
if value is None:
|
||||
return "None"
|
||||
elif isinstance(value, torch.Tensor):
|
||||
return f"tensor(shape={list(value.shape)}, dtype={value.dtype})"
|
||||
elif isinstance(value, list):
|
||||
if len(value) == 0:
|
||||
return "[]"
|
||||
elif len(value) <= 3:
|
||||
item_strs = [self._format_value(item) for item in value]
|
||||
return f"[{', '.join(item_strs)}]"
|
||||
else:
|
||||
return f"list(length={len(value)}, first_item={self._format_value(value[0])})"
|
||||
elif isinstance(value, str):
|
||||
if len(value) > 50:
|
||||
return f"'{value[:47]}...'"
|
||||
else:
|
||||
return f"'{value}'"
|
||||
else:
|
||||
return f"{type(value).__name__}({value})"
|
||||
|
||||
|
||||
class VerificationResult:
|
||||
"""Wrapper class for stage verification results."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._checks: Dict[str, bool] = {}
|
||||
self._failures: Dict[str, List[ValidationFailure]] = {}
|
||||
|
||||
def add_check(
|
||||
self, field_name: str, value: Any,
|
||||
validators: Union[Callable[[Any], bool], List[Callable[[Any], bool]]]
|
||||
) -> 'VerificationResult':
|
||||
"""
|
||||
Add a validation check for a field.
|
||||
|
||||
Args:
|
||||
field_name: Name of the field being checked
|
||||
value: The actual value to validate
|
||||
validators: Single validation function or list of validation functions.
|
||||
Each function will be called with the value as its first argument.
|
||||
|
||||
Returns:
|
||||
Self for method chaining
|
||||
|
||||
Examples:
|
||||
# Single validator
|
||||
result.add_check("tensor", my_tensor, V.is_tensor)
|
||||
|
||||
# Multiple validators (all must pass)
|
||||
result.add_check("latents", batch.latents, [V.is_tensor, V.with_dims(5)])
|
||||
|
||||
# Using partial functions for parameters
|
||||
result.add_check("height", batch.height, [V.not_none, V.divisible(8)])
|
||||
"""
|
||||
if not isinstance(validators, list):
|
||||
validators = [validators]
|
||||
|
||||
failures = []
|
||||
all_passed = True
|
||||
|
||||
# Apply all validators and collect detailed failure info
|
||||
for validator in validators:
|
||||
try:
|
||||
passed = validator(value)
|
||||
if not passed:
|
||||
all_passed = False
|
||||
failure = self._create_validation_failure(validator, value)
|
||||
failures.append(failure)
|
||||
except Exception as e:
|
||||
# If any validator raises an exception, consider the check failed
|
||||
all_passed = False
|
||||
validator_name = getattr(validator, '__name__', str(validator))
|
||||
failure = ValidationFailure(
|
||||
validator_name=validator_name,
|
||||
actual_value=value,
|
||||
error_msg=f"Exception during validation: {str(e)}")
|
||||
failures.append(failure)
|
||||
|
||||
self._checks[field_name] = all_passed
|
||||
if not all_passed:
|
||||
self._failures[field_name] = failures
|
||||
|
||||
return self
|
||||
|
||||
def _create_validation_failure(self, validator: Callable,
|
||||
value: Any) -> ValidationFailure:
|
||||
"""Create a ValidationFailure with detailed information."""
|
||||
validator_name = getattr(validator, '__name__', str(validator))
|
||||
|
||||
# Try to extract meaningful expected value info based on validator type
|
||||
expected = None
|
||||
error_msg = None
|
||||
|
||||
# Handle common validator patterns
|
||||
if hasattr(validator, '__closure__') and validator.__closure__:
|
||||
# This is likely a closure (like our helper functions)
|
||||
if 'dims' in validator_name or 'with_dims' in str(validator):
|
||||
if isinstance(value, torch.Tensor):
|
||||
expected = f"tensor with {validator.__closure__[0].cell_contents} dimensions"
|
||||
else:
|
||||
expected = "tensor with specific dimensions"
|
||||
elif 'divisible' in str(validator):
|
||||
expected = f"integer divisible by {validator.__closure__[0].cell_contents}"
|
||||
|
||||
# Handle specific validator types and check for NaN values
|
||||
if validator_name == 'is_tensor':
|
||||
expected = "torch.Tensor without NaN values"
|
||||
if isinstance(value,
|
||||
torch.Tensor) and torch.isnan(value).any().item():
|
||||
error_msg = f"tensor contains {torch.isnan(value).sum().item()} NaN values"
|
||||
elif validator_name == 'positive_int':
|
||||
expected = "positive integer"
|
||||
elif validator_name == 'not_none':
|
||||
expected = "non-None value"
|
||||
elif validator_name == 'list_not_empty':
|
||||
expected = "non-empty list"
|
||||
elif validator_name == 'bool_value':
|
||||
expected = "boolean value"
|
||||
elif 'tensor_with_dims' in validator_name or 'tensor_min_dims' in validator_name:
|
||||
if isinstance(value, torch.Tensor):
|
||||
if torch.isnan(value).any().item():
|
||||
error_msg = f"tensor has {value.dim()} dimensions but contains {torch.isnan(value).sum().item()} NaN values"
|
||||
else:
|
||||
error_msg = f"tensor has {value.dim()} dimensions"
|
||||
elif validator_name == 'is_list':
|
||||
expected = "list"
|
||||
elif validator_name == 'none_or_tensor':
|
||||
expected = "None or tensor without NaN values"
|
||||
if isinstance(value,
|
||||
torch.Tensor) and torch.isnan(value).any().item():
|
||||
error_msg = f"tensor contains {torch.isnan(value).sum().item()} NaN values"
|
||||
elif validator_name == 'list_of_tensors':
|
||||
expected = "non-empty list of tensors without NaN values"
|
||||
if isinstance(value, list) and len(value) > 0:
|
||||
nan_count = 0
|
||||
for item in value:
|
||||
if isinstance(
|
||||
item,
|
||||
torch.Tensor) and torch.isnan(item).any().item():
|
||||
nan_count += torch.isnan(item).sum().item()
|
||||
if nan_count > 0:
|
||||
error_msg = f"list contains tensors with total {nan_count} NaN values"
|
||||
elif 'list_of_tensors_with_dims' in validator_name:
|
||||
expected = "non-empty list of tensors with specific dimensions and no NaN values"
|
||||
if isinstance(value, list) and len(value) > 0:
|
||||
nan_count = 0
|
||||
for item in value:
|
||||
if isinstance(
|
||||
item,
|
||||
torch.Tensor) and torch.isnan(item).any().item():
|
||||
nan_count += torch.isnan(item).sum().item()
|
||||
if nan_count > 0:
|
||||
error_msg = f"list contains tensors with total {nan_count} NaN values"
|
||||
|
||||
return ValidationFailure(validator_name=validator_name,
|
||||
actual_value=value,
|
||||
expected=expected,
|
||||
error_msg=error_msg)
|
||||
|
||||
def is_valid(self) -> bool:
|
||||
"""Check if all validations passed."""
|
||||
return all(self._checks.values())
|
||||
|
||||
def get_failed_fields(self) -> List[str]:
|
||||
"""Get list of fields that failed validation."""
|
||||
return [field for field, passed in self._checks.items() if not passed]
|
||||
|
||||
def get_detailed_failures(self) -> Dict[str, List[ValidationFailure]]:
|
||||
"""Get detailed failure information for each failed field."""
|
||||
return self._failures.copy()
|
||||
|
||||
def get_failure_summary(self) -> str:
|
||||
"""Get a comprehensive summary of all validation failures."""
|
||||
if self.is_valid():
|
||||
return "All validations passed"
|
||||
|
||||
summary_parts = []
|
||||
for field_name, failures in self._failures.items():
|
||||
field_summary = f"\n Field '{field_name}':"
|
||||
for i, failure in enumerate(failures, 1):
|
||||
field_summary += f"\n {i}. {failure}"
|
||||
summary_parts.append(field_summary)
|
||||
|
||||
return "Validation failures:" + "".join(summary_parts)
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""Convert to dictionary for backward compatibility."""
|
||||
return self._checks.copy()
|
||||
|
||||
|
||||
# Alias for convenience
|
||||
V = StageValidators
|
||||
@@ -76,4 +76,36 @@ class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
|
||||
class WanImageToVideoValidationPipeline(ComposedPipelineBase):
|
||||
"""
|
||||
I2V Validation pipeline for Wan2.1, assumes that the input are preprocess latents.
|
||||
"""
|
||||
_required_config_modules = ["vae", "scheduler", "transformer"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer")))
|
||||
|
||||
self.add_stage(stage_name="image_latent_preparation_stage",
|
||||
stage=EncodingStage(vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
|
||||
EntryClass = WanImageToVideoPipeline
|
||||
|
||||
@@ -123,14 +123,13 @@ class CudaPlatformBase(Platform):
|
||||
SlidingTileAttentionBackend)
|
||||
logger.info("Using Sliding Tile Attention backend.")
|
||||
|
||||
# Overwrite with the actual backend
|
||||
envs.FASTVIDEO_ATTENTION_BACKEND = "SLIDING_TILE_ATTN"
|
||||
return "fastvideo.v1.attention.backends.sliding_tile_attn.SlidingTileAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
logger.info(
|
||||
"Sliding Tile Attention backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
logger.error(
|
||||
"Failed to import Sliding Tile Attention backend: %s",
|
||||
str(e))
|
||||
raise ImportError(
|
||||
"Sliding Tile Attention backend is not installed. ") from e
|
||||
elif selected_backend == AttentionBackendEnum.SAGE_ATTN:
|
||||
try:
|
||||
from sageattention import sageattn # noqa: F401
|
||||
@@ -139,8 +138,6 @@ class CudaPlatformBase(Platform):
|
||||
SageAttentionBackend)
|
||||
logger.info("Using Sage Attention backend.")
|
||||
|
||||
# Overwrite with the actual backend
|
||||
envs.FASTVIDEO_ATTENTION_BACKEND = "SAGE_ATTN"
|
||||
return "fastvideo.v1.attention.backends.sage_attn.SageAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
@@ -155,14 +152,13 @@ class CudaPlatformBase(Platform):
|
||||
VideoSparseAttentionBackend)
|
||||
logger.info("Using Video Sparse Attention backend.")
|
||||
|
||||
# Overwrite with the actual backend
|
||||
envs.FASTVIDEO_ATTENTION_BACKEND = "VIDEO_SPARSE_ATTN"
|
||||
return "fastvideo.v1.attention.backends.video_sparse_attn.VideoSparseAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
logger.info(
|
||||
"Video Sparse Attention backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
logger.error(
|
||||
"Failed to import Video Sparse Attention backend: %s",
|
||||
str(e))
|
||||
raise ImportError(
|
||||
"Video Sparse Attention backend is not installed. ") from e
|
||||
elif selected_backend == AttentionBackendEnum.TORCH_SDPA:
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
|
||||
@@ -209,14 +205,10 @@ class CudaPlatformBase(Platform):
|
||||
if target_backend == AttentionBackendEnum.TORCH_SDPA:
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
|
||||
# Overwrite with the actual backend
|
||||
envs.FASTVIDEO_ATTENTION_BACKEND = "TORCH_SDPA"
|
||||
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
|
||||
|
||||
logger.info("Using Flash Attention backend.")
|
||||
|
||||
# Overwrite with the actual backend
|
||||
envs.FASTVIDEO_ATTENTION_BACKEND = "FLASH_ATTN"
|
||||
return "fastvideo.v1.attention.backends.flash_attn.FlashAttentionBackend"
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -169,6 +169,7 @@ class Platform:
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
|
||||
@classmethod
|
||||
def verify_model_arch(cls, model_arch: str) -> None:
|
||||
|
||||
@@ -168,5 +168,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()}"
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
import os
|
||||
import sys
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
import pytest
|
||||
|
||||
NUM_NODES = "1"
|
||||
NUM_GPUS_PER_NODE = "2"
|
||||
|
||||
# Set environment variables
|
||||
os.environ["FASTVIDEO_ATTENTION_CONFIG"] = "assets/mask_strategy_wan.json"
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLIDING_TILE_ATTN"
|
||||
|
||||
def test_inference():
|
||||
"""Test the inference functionality"""
|
||||
# Create command as in wan_14B-STA.sh
|
||||
cmd = [
|
||||
"fastvideo", "generate",
|
||||
"--model-path", "Wan-AI/Wan2.1-T2V-14B-Diffusers",
|
||||
"--sp-size", "2",
|
||||
"--tp-size", "2",
|
||||
"--num-gpus", "2",
|
||||
"--height", "768",
|
||||
"--width", "1280",
|
||||
"--num-frames", "69",
|
||||
"--num-inference-steps", "2",
|
||||
"--fps", "16",
|
||||
"--guidance-scale", "5.0",
|
||||
"--flow-shift", "5.0",
|
||||
"--prompt", "A majestic lion strides across the golden savanna, its powerful frame glistening under the warm afternoon sun. The tall grass ripples gently in the breeze, enhancing the lion's commanding presence. The tone is vibrant, embodying the raw energy of the wild. Low angle, steady tracking shot, cinematic.",
|
||||
"--negative-prompt", "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
"--seed", "1024",
|
||||
"--output-path", "outputs_video/STA_1024/",
|
||||
]
|
||||
|
||||
# Run the command
|
||||
subprocess.run(cmd, check=True)
|
||||
|
||||
# Verify output directory exists
|
||||
output_dir = Path("outputs_video/STA_1024/")
|
||||
assert output_dir.exists(), f"Output directory {output_dir} does not exist"
|
||||
|
||||
# Verify that video files were generated
|
||||
video_files = list(output_dir.glob("*.mp4"))
|
||||
assert len(video_files) > 0, "No video files were generated"
|
||||
|
||||
# Verify the video file properties
|
||||
for video_file in video_files:
|
||||
assert video_file.stat().st_size > 0, f"Video file {video_file} is empty"
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_inference()
|
||||
@@ -0,0 +1,99 @@
|
||||
import modal
|
||||
|
||||
app = modal.App()
|
||||
|
||||
import os
|
||||
|
||||
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",
|
||||
"BUILDKITE_REPO": os.environ.get("BUILDKITE_REPO", ""),
|
||||
"BUILDKITE_COMMIT": os.environ.get("BUILDKITE_COMMIT", ""),
|
||||
"BUILDKITE_PULL_REQUEST": os.environ.get("BUILDKITE_PULL_REQUEST", ""),
|
||||
"IMAGE_VERSION": os.environ.get("IMAGE_VERSION", ""),
|
||||
})
|
||||
)
|
||||
|
||||
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")
|
||||
pr_number = os.environ.get("BUILDKITE_PULL_REQUEST")
|
||||
|
||||
print(f"Cloning repository: {git_repo}")
|
||||
print(f"Target commit: {git_commit}")
|
||||
if pr_number:
|
||||
print(f"PR number: {pr_number}")
|
||||
|
||||
# For PRs (including forks), use GitHub's PR refs to get the correct commit
|
||||
if pr_number and pr_number != "false":
|
||||
checkout_command = f"git fetch --prune origin refs/pull/{pr_number}/head && git checkout FETCH_HEAD"
|
||||
print(f"Using PR ref for checkout: {checkout_command}")
|
||||
else:
|
||||
checkout_command = f"git checkout {git_commit}"
|
||||
print(f"Using direct commit checkout: {checkout_command}")
|
||||
|
||||
command = f"""
|
||||
source $HOME/.local/bin/env &&
|
||||
source /opt/venv/bin/activate &&
|
||||
git clone {git_repo} /FastVideo &&
|
||||
cd /FastVideo &&
|
||||
{checkout_command} &&
|
||||
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=900)
|
||||
def run_encoder_tests():
|
||||
run_test("pytest ./fastvideo/v1/tests/encoders -vs")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
def run_vae_tests():
|
||||
run_test("pytest ./fastvideo/v1/tests/vaes -vs")
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=900)
|
||||
def run_transformer_tests():
|
||||
run_test("pytest ./fastvideo/v1/tests/transformers -vs")
|
||||
|
||||
@app.function(gpu="L40S:2", image=image, timeout=1800)
|
||||
def run_ssim_tests():
|
||||
run_test("pytest ./fastvideo/v1/tests/ssim -vs")
|
||||
|
||||
@app.function(gpu="L40S:4", image=image, timeout=900, 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:2", image=image, timeout=900, 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:2", image=image, timeout=900)
|
||||
def run_inference_tests_STA():
|
||||
run_test("pytest ./fastvideo/v1/tests/inference/STA -srP")
|
||||
|
||||
@app.function(gpu="H100:1", image=image, timeout=900)
|
||||
def run_precision_tests_STA():
|
||||
run_test("python csrc/attn/tests/test_sta.py")
|
||||
|
||||
@app.function(gpu="H100:1", image=image, timeout=900)
|
||||
def run_precision_tests_VSA():
|
||||
run_test("python csrc/attn/tests/test_block_sparse.py")
|
||||
@@ -0,0 +1 @@
|
||||
{"step_time":2.245914653001819,"_wandb":{"runtime":1434},"learning_rate":1e-05,"grad_norm":0.57421875,"avg_step_time":1.1814782944297622,"train_loss":0.07932619750499725,"vsa_sparsity":0,"_timestamp":1.750578625921253e+09,"validation_videos_40_steps":{"count":1,"videos":[{"size":420969,"path":"media/videos/validation_videos_40_steps_900_581ff5eae2909d3a7b36.mp4","_type":"video-file","sha256":"581ff5eae2909d3a7b362dcb24d060c006c09e4d4deb44b82f4aa697f6789ba7"}],"captions":false,"_type":"videos"},"_runtime":1434.62395329,"_step":901}
|
||||
BIN
Binary file not shown.
@@ -0,0 +1,177 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
from huggingface_hub import snapshot_download
|
||||
import shutil
|
||||
import subprocess
|
||||
import sys
|
||||
from fastvideo.v1.tests.ssim.test_inference_similarity import compute_video_ssim_torchvision
|
||||
|
||||
# Import the training pipeline
|
||||
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
|
||||
|
||||
NUM_NODES = "1"
|
||||
MODEL_PATH = "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
|
||||
# preprocessing
|
||||
DATA_DIR = "data"
|
||||
LOCAL_RAW_DATA_DIR = Path(os.path.join(DATA_DIR, "cats"))
|
||||
NUM_GPUS_PER_NODE_PREPROCESSING = "1"
|
||||
PREPROCESSING_ENTRY_FILE_PATH = "fastvideo/v1/pipelines/preprocess/v1_preprocess.py"
|
||||
|
||||
LOCAL_PREPROCESSED_DATA_DIR = Path(os.path.join(DATA_DIR, "cats_preprocessed_data_i2v"))
|
||||
|
||||
|
||||
# training
|
||||
NUM_GPUS_PER_NODE_TRAINING = "8"
|
||||
TRAINING_ENTRY_FILE_PATH = "fastvideo/v1/training/wan_i2v_training_pipeline.py"
|
||||
LOCAL_TRAINING_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "combined_parquet_dataset")
|
||||
LOCAL_VALIDATION_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "validation_parquet_dataset")
|
||||
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)
|
||||
|
||||
print(f"Downloading raw dataset to {LOCAL_RAW_DATA_DIR}...")
|
||||
try:
|
||||
# result = snapshot_download(
|
||||
# repo_id="wlsaidhi/cats-overfit-merged",
|
||||
# local_dir=str(LOCAL_RAW_DATA_DIR),
|
||||
# repo_type="dataset",
|
||||
# resume_download=True,
|
||||
# token=os.environ.get("HF_TOKEN"), # In case authentication is needed
|
||||
# )
|
||||
print(f"Download completed successfully. Files downloaded to: {result}")
|
||||
|
||||
# Verify the download
|
||||
if not LOCAL_RAW_DATA_DIR.exists():
|
||||
raise RuntimeError(f"Download appeared to succeed but {LOCAL_RAW_DATA_DIR} does not exist")
|
||||
|
||||
# List downloaded files
|
||||
print("Downloaded files:")
|
||||
for file in LOCAL_RAW_DATA_DIR.rglob("*"):
|
||||
if file.is_file():
|
||||
print(f" - {file.relative_to(LOCAL_RAW_DATA_DIR)}")
|
||||
|
||||
except Exception as e:
|
||||
print(f"Error during download: {str(e)}")
|
||||
raise
|
||||
|
||||
|
||||
def run_preprocessing():
|
||||
# Run torchrun command
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE_PREPROCESSING,
|
||||
PREPROCESSING_ENTRY_FILE_PATH,
|
||||
"--model_path", MODEL_PATH,
|
||||
"--data_merge_path", os.path.join(LOCAL_RAW_DATA_DIR, "merge_1_sample.txt"),
|
||||
"--preprocess_video_batch_size", "1",
|
||||
"--max_height", "480",
|
||||
"--max_width", "832",
|
||||
"--num_frames", "77",
|
||||
"--dataloader_num_workers", "0",
|
||||
"--output_dir", LOCAL_PREPROCESSED_DATA_DIR,
|
||||
"--train_fps", "16",
|
||||
"--validation_dataset_file", os.path.join(LOCAL_RAW_DATA_DIR, "validation_i2v_prompt_1_sample.json"),
|
||||
"--samples_per_file", "1",
|
||||
"--flush_frequency", "1",
|
||||
"--video_length_tolerance_range", "5",
|
||||
"--preprocess_task", "i2v",
|
||||
]
|
||||
|
||||
process = subprocess.run(cmd, check=True)
|
||||
|
||||
|
||||
def run_training():
|
||||
cmd = [
|
||||
"torchrun",
|
||||
"--nnodes", NUM_NODES,
|
||||
"--nproc_per_node", NUM_GPUS_PER_NODE_TRAINING,
|
||||
TRAINING_ENTRY_FILE_PATH,
|
||||
"--model_path", MODEL_PATH,
|
||||
"--inference_mode", "False",
|
||||
"--pretrained_model_name_or_path", MODEL_PATH,
|
||||
"--data_path", LOCAL_TRAINING_DATA_DIR,
|
||||
"--validation_preprocessed_path", LOCAL_VALIDATION_DATA_DIR,
|
||||
"--train_batch_size", "1",
|
||||
"--num_latent_t", "8",
|
||||
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--sp_size", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--tp_size", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--hsdp_replicate_dim", "1",
|
||||
"--hsdp_shard_dim", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
|
||||
"--train_sp_batch_size", "1",
|
||||
"--dataloader_num_workers", "10",
|
||||
"--gradient_accumulation_steps", "1",
|
||||
"--max_train_steps", "901",
|
||||
"--learning_rate", "1e-5",
|
||||
"--mixed_precision", "bf16",
|
||||
"--checkpointing_steps", "6000",
|
||||
"--validation_steps", "100",
|
||||
"--validation_sampling_steps", "40",
|
||||
"--log_validation",
|
||||
"--checkpoints_total_limit", "3",
|
||||
"--allow_tf32",
|
||||
"--ema_start_step", "0",
|
||||
"--training_cfg_rate", "0.1",
|
||||
"--output_dir", LOCAL_OUTPUT_DIR,
|
||||
"--tracker_project_name", "wan_i2v_finetune_overfit_ci",
|
||||
"--num_height", "480",
|
||||
"--num_width", "832",
|
||||
"--num_frames", "81",
|
||||
"--validation_guidance_scale", "1.0",
|
||||
"--num_euler_timesteps", "50",
|
||||
"--multi_phased_distill_schedule", "4000-1",
|
||||
"--weight_decay", "0.01",
|
||||
"--not_apply_cfg_solver",
|
||||
"--dit_precision", "fp32",
|
||||
"--max_grad_norm", "1.0",
|
||||
]
|
||||
|
||||
print(f"Running training with command: {cmd}")
|
||||
process = subprocess.run(cmd, check=True)
|
||||
|
||||
|
||||
def test_e2e_overfit_single_sample():
|
||||
os.environ["WANDB_MODE"] = "online"
|
||||
|
||||
# download_data()
|
||||
run_preprocessing()
|
||||
run_training()
|
||||
|
||||
reference_video_file = os.path.join(os.path.dirname(__file__), "reference_video_1_sample_v0.mp4")
|
||||
print(f"reference_video_file: {reference_video_file}")
|
||||
final_validation_video_file = os.path.join(LOCAL_OUTPUT_DIR, "validation_step_900_inference_steps_50_video_0.mp4")
|
||||
print(f"final_validation_video_file: {final_validation_video_file}")
|
||||
|
||||
|
||||
# Ensure both files exist
|
||||
assert os.path.exists(reference_video_file), f"Reference video not found at {reference_video_file}"
|
||||
assert os.path.exists(final_validation_video_file), f"Validation video not found at {final_validation_video_file}"
|
||||
|
||||
# Compute SSIM
|
||||
mean_ssim, min_ssim, max_ssim = compute_video_ssim_torchvision(
|
||||
reference_video_file,
|
||||
final_validation_video_file,
|
||||
use_ms_ssim=True # Using MS-SSIM for better quality assessment
|
||||
)
|
||||
|
||||
print("\n===== SSIM Results for Step 900 Validation =====")
|
||||
print(f"Mean MS-SSIM: {mean_ssim:.4f}")
|
||||
print(f"Min MS-SSIM: {min_ssim:.4f}")
|
||||
print(f"Max MS-SSIM: {max_ssim:.4f}")
|
||||
|
||||
assert max_ssim > 0.5, f"Max SSIM is below 0.5: {max_ssim}"
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_e2e_overfit_single_sample()
|
||||
@@ -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:
|
||||
@@ -80,7 +77,7 @@ def run_preprocessing():
|
||||
"--dataloader_num_workers", "0",
|
||||
"--output_dir", LOCAL_PREPROCESSED_DATA_DIR,
|
||||
"--train_fps", "16",
|
||||
"--validation_prompt_txt", os.path.join(LOCAL_RAW_DATA_DIR, "validation_prompt_1_sample.txt"),
|
||||
"--validation_dataset_file", os.path.join(LOCAL_RAW_DATA_DIR, "validation_prompt_1_sample.json"),
|
||||
"--samples_per_file", "1",
|
||||
"--flush_frequency", "1",
|
||||
"--video_length_tolerance_range", "5",
|
||||
@@ -100,7 +97,7 @@ def run_training():
|
||||
"--inference_mode", "False",
|
||||
"--pretrained_model_name_or_path", MODEL_PATH,
|
||||
"--data_path", LOCAL_TRAINING_DATA_DIR,
|
||||
"--validation_prompt_dir", LOCAL_VALIDATION_DATA_DIR,
|
||||
"--validation_preprocessed_path", LOCAL_VALIDATION_DATA_DIR,
|
||||
"--train_batch_size", "1",
|
||||
"--num_latent_t", "8",
|
||||
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
|
||||
@@ -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.
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user