Compare commits
25
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
480868bef9 | ||
|
|
aab74c1271 | ||
|
|
f89d86944f | ||
|
|
8741d204a5 | ||
|
|
cdc85f58a8 | ||
|
|
0262d2f089 | ||
|
|
62c0343465 | ||
|
|
1e1a023fb0 | ||
|
|
1d2517ad8e | ||
|
|
d41186cb4a | ||
|
|
78e0c7eec9 | ||
|
|
1c41a94b62 | ||
|
|
2e66aafe20 | ||
|
|
55074bda76 | ||
|
|
de65bec2b7 | ||
|
|
7664dd0de3 | ||
|
|
019a88ced4 | ||
|
|
72de11abcc | ||
|
|
d71a4ebffc | ||
|
|
1089ab43bf | ||
|
|
97d4b984c9 | ||
|
|
2a8953d74d | ||
|
|
8801b10da7 | ||
|
|
6b413f2ec4 | ||
|
|
28b72694aa |
@@ -0,0 +1,66 @@
|
||||
env:
|
||||
IMAGE_VERSION: "py3.12-latest"
|
||||
|
||||
steps:
|
||||
- block: "Start Build"
|
||||
blocked_state: "running"
|
||||
prompt: "Approve build?"
|
||||
|
||||
- label: "Trigger Tests"
|
||||
command: |
|
||||
echo "Current working directory: $(pwd)"
|
||||
echo "Current branch:"
|
||||
git branch --show-current
|
||||
echo "Full diff:"
|
||||
git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD
|
||||
plugins:
|
||||
- monorepo-diff#v1.4.0:
|
||||
diff: "git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD"
|
||||
watch:
|
||||
- path:
|
||||
- "fastvideo/v1/models/encoders/**"
|
||||
- "fastvideo/v1/models/loaders/**"
|
||||
- "fastvideo/v1/tests/encoders/**"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Encoder Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=encoder
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/models/vaes/**"
|
||||
- "fastvideo/v1/models/loaders/**"
|
||||
- "fastvideo/v1/tests/vaes/**"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "VAE Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=vae
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/models/dits/**"
|
||||
- "fastvideo/v1/models/loaders/**"
|
||||
- "fastvideo/v1/tests/transformers/**"
|
||||
- "fastvideo/v1/layers/**"
|
||||
- "fastvideo/v1/attention/**"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Transformer Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=transformer
|
||||
agents:
|
||||
queue: "default"
|
||||
- path: "fastvideo/v1/**/*.py"
|
||||
config:
|
||||
command: "timeout 60m .buildkite/scripts/pr_test.sh"
|
||||
label: "SSIM Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=ssim
|
||||
agents:
|
||||
queue: "default"
|
||||
Executable
+91
@@ -0,0 +1,91 @@
|
||||
#!/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)
|
||||
|
||||
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"
|
||||
|
||||
case "$TEST_TYPE" in
|
||||
"encoder")
|
||||
log "Running encoder tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
|
||||
;;
|
||||
"vae")
|
||||
log "Running VAE tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
|
||||
;;
|
||||
"transformer")
|
||||
log "Running transformer tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
|
||||
;;
|
||||
"ssim")
|
||||
log "Running SSIM tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
;;
|
||||
*)
|
||||
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
|
||||
@@ -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
|
||||
]
|
||||
|
||||
|
||||
+180
-16
@@ -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,10 +42,36 @@ 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
|
||||
default: false
|
||||
type: boolean
|
||||
|
||||
env:
|
||||
PYTHONUNBUFFERED: "1"
|
||||
|
||||
|
||||
concurrency:
|
||||
group: pr-test-${{ github.ref }}
|
||||
cancel-in-progress: true
|
||||
@@ -65,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
|
||||
@@ -99,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 }}
|
||||
@@ -117,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 }}
|
||||
@@ -135,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 }}
|
||||
@@ -163,7 +230,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 }}
|
||||
@@ -181,15 +248,112 @@ 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' && github.event.pull_request.draft == false) ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test_VSA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "training-test-VSA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/VSA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
inference-test-STA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_inference_test_STA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "inference-test-STA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/inference/STA -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
precision-test-STA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(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' && github.event.pull_request.draft == false) ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_VSA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "precision-test-VSA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_block_sparse.py"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
nightly-test:
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "nightly-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
|
||||
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/nightly/test_e2e_overfit_single_sample.py -vs"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
runpod-cleanup:
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
|
||||
# Add other jobs to this list as you create them
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test, training-test, training-test-VSA, inference-test-STA, precision-test-STA, precision-test-VSA]
|
||||
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
@@ -206,7 +370,7 @@ jobs:
|
||||
|
||||
- name: Cleanup all RunPod instances
|
||||
env:
|
||||
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12"]'
|
||||
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12", "training-test", "training-test-VSA", "inference-test-STA", "precision-test-STA", "precision-test-VSA"]'
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
run: python .github/scripts/runpod_cleanup.py
|
||||
|
||||
@@ -43,6 +43,8 @@ on:
|
||||
required: true
|
||||
RUNPOD_PRIVATE_KEY:
|
||||
required: true
|
||||
WANDB_API_KEY:
|
||||
required: false
|
||||
|
||||
jobs:
|
||||
run-test:
|
||||
@@ -55,7 +57,7 @@ jobs:
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Set up SSH key
|
||||
run: |
|
||||
@@ -72,6 +74,7 @@ jobs:
|
||||
JOB_ID: ${{ inputs.job_id }}
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
timeout-minutes: ${{ inputs.timeout_minutes }}
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
|
||||
@@ -91,7 +91,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
|
||||
|
||||
## Distillation and Finetuning
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/training/distillation.html)
|
||||
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetuning.html)
|
||||
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html)
|
||||
|
||||
## 📑 Development Plan
|
||||
|
||||
@@ -111,7 +111,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/developer_guide/overview.html)
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview.html)
|
||||
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects:
|
||||
|
||||
+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
|
||||
@@ -288,7 +288,7 @@ Sequence parallelism splits sequences across devices:
|
||||
|
||||
```python
|
||||
# Distributed attention for long sequences
|
||||
from fastvideo.v1.layers.attention import DistributedAttention
|
||||
from fastvideo.v1.attention import DistributedAttention
|
||||
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
|
||||
@@ -96,8 +96,8 @@ Replace standard attention with FastVideo's optimized attention:
|
||||
|
||||
```python
|
||||
# Local attention patterns
|
||||
from fastvideo.v1.layers.attention import LocalAttention
|
||||
from fastvideo.v1.layers.attention.backends.abstract import _Backend
|
||||
from fastvideo.v1.attention import LocalAttention
|
||||
from fastvideo.v1.attention.backends.abstract import _Backend
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
@@ -108,7 +108,7 @@ self.attn = LocalAttention(
|
||||
)
|
||||
|
||||
# Distributed attention for long sequences
|
||||
from fastvideo.v1.layers.attention import DistributedAttention
|
||||
from fastvideo.v1.attention import DistributedAttention
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
|
||||
@@ -13,6 +13,5 @@ This is is a demo for 2D STA with window size (6,6) operating on a (10, 10) imag
|
||||
You can run STA using the following command:
|
||||
|
||||
```bash
|
||||
huggingface-cli download hunyuanvideo-community/HunyuanVideo --local-dir data/hunyuan
|
||||
bash scripts/inference/inference_hunyuan_STA.sh
|
||||
bash scripts/inference/v1_inference_wan_STA.sh
|
||||
```
|
||||
|
||||
@@ -7,70 +7,40 @@ To save GPU memory, we precompute text embeddings and VAE latents to eliminate t
|
||||
We provide a sample dataset to help you get started. Download the source media using the following command:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Image-Vid-Finetune-Src --local_dir=data/Image-Vid-Finetune-Src --repo_type=dataset
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/mini_i2v_dataset --local_dir=FastVideo/mini_i2v_dataset --repo_type=dataset
|
||||
```
|
||||
|
||||
The folder `crush-smol_raw/` contains raw videos and captions for testing preprocessing, while `crush-smol_preprocessed/` contains latents prepared for testing training.
|
||||
|
||||
To preprocess the dataset for fine-tuning or distillation, run:
|
||||
|
||||
```
|
||||
bash scripts/preprocess/preprocess_mochi_data.sh # for mochi
|
||||
bash scripts/preprocess/preprocess_hunyuan_data.sh # for hunyuan
|
||||
bash scripts/preprocess/v1_preprocess_wan_data_t2v # for wan
|
||||
```
|
||||
|
||||
The preprocessed dataset will be stored in `Image-Vid-Finetune-Mochi` or `Image-Vid-Finetune-HunYuan` correspondingly.
|
||||
|
||||
## Process your own dataset
|
||||
|
||||
If you wish to create your own dataset for finetuning or distillation, please structure you video dataset in the following format:
|
||||
If you wish to create your own dataset for finetuning or distillation, please refer `mini_i2v_dataset/crush-smol_raw/` to structure you video dataset in the following format:
|
||||
|
||||
```
|
||||
path_to_dataset_folder/
|
||||
├── media/
|
||||
│ ├── 0.jpg
|
||||
path_to_your_dataset_folder/
|
||||
├── videos/
|
||||
│ ├── 0.mp4
|
||||
│ ├── 1.mp4
|
||||
│ ├── 2.jpg
|
||||
├── video2caption.json
|
||||
└── merge.txt
|
||||
├── videos.txt
|
||||
└── prompt.txt
|
||||
```
|
||||
|
||||
Format the JSON file as a list, where each item represents a media source:
|
||||
To geranate the `videos2caption.json` and `merge.txt`, run
|
||||
|
||||
For image media,
|
||||
|
||||
```
|
||||
{
|
||||
"path": "0.jpg",
|
||||
"cap": ["captions"]
|
||||
}
|
||||
``` python
|
||||
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
|
||||
```
|
||||
|
||||
For video media,
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/v1_preprocess_****.sh` accordingly and run:
|
||||
|
||||
```
|
||||
{
|
||||
"path": "1.mp4",
|
||||
"resolution": {
|
||||
"width": 848,
|
||||
"height": 480
|
||||
},
|
||||
"fps": 30.0,
|
||||
"duration": 6.033333333333333,
|
||||
"cap": [
|
||||
"caption"
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
Use a txt file (merge.txt) to contain the source folder for media and the JSON file for meta information:
|
||||
|
||||
```
|
||||
path_to_media_source_foder,path_to_json_file
|
||||
```
|
||||
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/preprocess_****_data.sh` accordingly and run:
|
||||
|
||||
```
|
||||
bash scripts/preprocess/preprocess_****_data.sh
|
||||
bash scripts/preprocess/v1_preprocess_****.sh
|
||||
```
|
||||
|
||||
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
|
||||
|
||||
@@ -16,6 +16,13 @@ bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
## ⚡ Finetune with VSA
|
||||
Follow [data_preprocess.md](#v0-data-preprocess) to get parquet files for preproccessed latent, and then run:
|
||||
|
||||
```bash
|
||||
bash scripts/finetune/finetune_v1_VSA.sh
|
||||
```
|
||||
|
||||
## ⚡ Lora Finetune
|
||||
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
|
||||
@@ -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,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": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-034.mp4",
|
||||
"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": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-027.mp4",
|
||||
"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": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-030.mp4",
|
||||
"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,
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.attention.layer import (DistributedAttention,
|
||||
DistributedAttention_VSA,
|
||||
LocalAttention)
|
||||
from fastvideo.v1.attention.selector import get_attn_backend
|
||||
|
||||
__all__ = [
|
||||
"DistributedAttention",
|
||||
"LocalAttention",
|
||||
"DistributedAttention_VSA",
|
||||
"AttentionBackend",
|
||||
"AttentionMetadata",
|
||||
"AttentionMetadataBuilder",
|
||||
# "AttentionState",
|
||||
"get_attn_backend",
|
||||
]
|
||||
+4
-3
@@ -14,9 +14,10 @@ try:
|
||||
except ImportError:
|
||||
flash_attn_func = flash_attn_2_func
|
||||
|
||||
from fastvideo.v1.layers.attention.backends.abstract import (
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
+3
-3
@@ -4,10 +4,10 @@ from typing import List, Optional, Type
|
||||
import torch
|
||||
from sageattention import sageattn
|
||||
|
||||
from fastvideo.v1.layers.attention.backends.abstract import (
|
||||
from fastvideo.v1.attention.backends.abstract import (
|
||||
AttentionBackend) # FlashAttentionMetadata,
|
||||
from fastvideo.v1.layers.attention.backends.abstract import (AttentionImpl,
|
||||
AttentionMetadata)
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionImpl,
|
||||
AttentionMetadata)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
+3
-3
@@ -3,10 +3,10 @@ from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.layers.attention.backends.abstract import (
|
||||
from fastvideo.v1.attention.backends.abstract import (
|
||||
AttentionBackend) # FlashAttentionMetadata,
|
||||
from fastvideo.v1.layers.attention.backends.abstract import (AttentionImpl,
|
||||
AttentionMetadata)
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionImpl,
|
||||
AttentionMetadata)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
+6
-27
@@ -1,48 +1,27 @@
|
||||
# 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
|
||||
from st_attn import sliding_tile_attention
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.distributed import get_sp_group
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.v1.layers.attention.backends.abstract import (
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
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:
|
||||
+4
-3
@@ -11,11 +11,12 @@ try:
|
||||
except ImportError:
|
||||
video_sparse_attn = None
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.distributed import get_sp_group
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.layers.attention.backends.abstract import (
|
||||
AttentionBackend, AttentionImpl, AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
@@ -5,14 +5,14 @@ from typing import Optional, Tuple
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.v1.attention.selector import (backend_name_to_enum,
|
||||
get_attn_backend)
|
||||
from fastvideo.v1.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather, sequence_model_parallel_all_to_all_4D)
|
||||
from fastvideo.v1.distributed.parallel_state import (get_sp_parallel_rank,
|
||||
get_sp_world_size)
|
||||
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.v1.layers.attention.selector import (backend_name_to_enum,
|
||||
get_attn_backend)
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
from fastvideo.v1.utils import get_compute_dtype
|
||||
|
||||
|
||||
@@ -26,8 +26,8 @@ class DistributedAttention(nn.Module):
|
||||
num_kv_heads: Optional[int] = None,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[
|
||||
AttentionBackendEnum, ...]] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
@@ -211,8 +211,8 @@ class LocalAttention(nn.Module):
|
||||
num_kv_heads: Optional[int] = None,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[
|
||||
AttentionBackendEnum, ...]] = None,
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
if softmax_scale is None:
|
||||
@@ -9,15 +9,15 @@ from typing import Generator, Optional, Tuple, Type, cast
|
||||
import torch
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.layers.attention.backends.abstract import AttentionBackend
|
||||
from fastvideo.v1.attention.backends.abstract import AttentionBackend
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.platforms import _Backend, current_platform
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum, current_platform
|
||||
from fastvideo.v1.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
|
||||
def backend_name_to_enum(backend_name: str) -> Optional[AttentionBackendEnum]:
|
||||
"""
|
||||
Convert a string backend name to a _Backend enum value.
|
||||
|
||||
@@ -27,11 +27,11 @@ def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
|
||||
loaded.
|
||||
"""
|
||||
assert backend_name is not None
|
||||
return _Backend[backend_name] if backend_name in _Backend.__members__ else \
|
||||
return AttentionBackendEnum[backend_name] if backend_name in AttentionBackendEnum.__members__ else \
|
||||
None
|
||||
|
||||
|
||||
def get_env_variable_attn_backend() -> Optional[_Backend]:
|
||||
def get_env_variable_attn_backend() -> Optional[AttentionBackendEnum]:
|
||||
'''
|
||||
Get the backend override specified by the FastVideo attention
|
||||
backend environment variable, if one is specified.
|
||||
@@ -53,10 +53,11 @@ def get_env_variable_attn_backend() -> Optional[_Backend]:
|
||||
#
|
||||
# THIS SELECTION TAKES PRECEDENCE OVER THE
|
||||
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
|
||||
forced_attn_backend: Optional[_Backend] = None
|
||||
forced_attn_backend: Optional[AttentionBackendEnum] = None
|
||||
|
||||
|
||||
def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
|
||||
def global_force_attn_backend(
|
||||
attn_backend: Optional[AttentionBackendEnum]) -> None:
|
||||
'''
|
||||
Force all attention operations to use a specified backend.
|
||||
|
||||
@@ -71,7 +72,7 @@ def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
|
||||
forced_attn_backend = attn_backend
|
||||
|
||||
|
||||
def get_global_forced_attn_backend() -> Optional[_Backend]:
|
||||
def get_global_forced_attn_backend() -> Optional[AttentionBackendEnum]:
|
||||
'''
|
||||
Get the currently-forced choice of attention backend,
|
||||
or None if auto-selection is currently enabled.
|
||||
@@ -82,7 +83,8 @@ def get_global_forced_attn_backend() -> Optional[_Backend]:
|
||||
def get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
) -> Type[AttentionBackend]:
|
||||
return _cached_get_attn_backend(head_size, dtype,
|
||||
supported_attention_backends)
|
||||
@@ -92,7 +94,8 @@ def get_attn_backend(
|
||||
def _cached_get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
) -> Type[AttentionBackend]:
|
||||
# Check whether a particular choice of backend was
|
||||
# previously forced.
|
||||
@@ -102,7 +105,7 @@ def _cached_get_attn_backend(
|
||||
if not supported_attention_backends:
|
||||
raise ValueError("supported_attention_backends is empty")
|
||||
selected_backend = None
|
||||
backend_by_global_setting: Optional[_Backend] = (
|
||||
backend_by_global_setting: Optional[AttentionBackendEnum] = (
|
||||
get_global_forced_attn_backend())
|
||||
if backend_by_global_setting is not None:
|
||||
selected_backend = backend_by_global_setting
|
||||
@@ -125,7 +128,7 @@ def _cached_get_attn_backend(
|
||||
|
||||
@contextmanager
|
||||
def global_force_attn_backend_context_manager(
|
||||
attn_backend: _Backend) -> Generator[None, None, None]:
|
||||
attn_backend: AttentionBackendEnum) -> Generator[None, None, None]:
|
||||
'''
|
||||
Globally force a FastVideo attention backend override within a
|
||||
context manager, reverting the global attention backend
|
||||
@@ -4,7 +4,7 @@ from typing import Any, List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -12,13 +12,12 @@ class DiTArchConfig(ArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=list)
|
||||
_compile_conditions: list = field(default_factory=list)
|
||||
_param_names_mapping: dict = field(default_factory=dict)
|
||||
_reverse_param_names_mapping: dict = field(default_factory=dict)
|
||||
_lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
_supported_attention_backends: Tuple[_Backend,
|
||||
...] = (_Backend.SLIDING_TILE_ATTN,
|
||||
_Backend.SAGE_ATTN,
|
||||
_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA,
|
||||
_Backend.VIDEO_SPARSE_ATTN)
|
||||
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN)
|
||||
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
|
||||
@@ -147,6 +147,9 @@ class HunyuanVideoArchConfig(DiTArchConfig):
|
||||
r"final_layer.linear.\1",
|
||||
})
|
||||
|
||||
# Reverse mapping for saving checkpoints: training -> diffusers
|
||||
_reverse_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
patch_size: int = 2
|
||||
patch_size_t: int = 1
|
||||
in_channels: int = 16
|
||||
|
||||
@@ -49,9 +49,13 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
r"blocks.\1.ffn.fc_in.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
|
||||
r"blocks.\1.ffn.fc_out.\2",
|
||||
r"blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"^blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"blocks.\1.self_attn_residual_norm.norm.\2",
|
||||
})
|
||||
|
||||
# Reverse mapping for saving checkpoints: training -> diffusers
|
||||
_reverse_param_names_mapping: dict = field(default_factory=lambda: {})
|
||||
|
||||
# Some LoRA adapters use the original official layer names instead of hf layer names,
|
||||
# so apply this before the param_names_mapping
|
||||
_lora_param_names_mapping: dict = field(
|
||||
|
||||
@@ -6,14 +6,14 @@ import torch
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
@dataclass
|
||||
class EncoderArchConfig(ArchConfig):
|
||||
architectures: List[str] = field(default_factory=lambda: [])
|
||||
_supported_attention_backends: Tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA)
|
||||
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA)
|
||||
output_hidden_states: bool = False
|
||||
use_return_dict: bool = True
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,17 @@ 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"
|
||||
]
|
||||
|
||||
@@ -48,6 +48,48 @@ pyarrow_schema_i2v = pa.schema([
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
pyarrow_schema_i2v_validation = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Image/Video VAE latents ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("vae_latent_bytes", pa.binary()),
|
||||
# e.g., [C, T, H, W] or [C, H, W]
|
||||
pa.field("vae_latent_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'float32'
|
||||
pa.field("vae_latent_dtype", 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()),
|
||||
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()),
|
||||
# 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_t2v = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Image/Video VAE latents ---
|
||||
@@ -80,4 +122,38 @@ 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()),
|
||||
# --- Image/Video VAE latents ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("vae_latent_bytes", pa.binary()),
|
||||
# e.g., [C, T, H, W] or [C, H, W]
|
||||
pa.field("vae_latent_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'float32'
|
||||
pa.field("vae_latent_dtype", 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()),
|
||||
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()),
|
||||
pa.field("media_type", pa.string()), # 'image' or 'video'
|
||||
pa.field("width", pa.int64()),
|
||||
pa.field("height", pa.int64()),
|
||||
# -- Video-specific (can be null/default for images) ---
|
||||
# Number of frames processed (e.g., 1 for image, N for video)
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
@@ -1,137 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from multiprocessing import Pool, cpu_count
|
||||
from pathlib import Path
|
||||
|
||||
import torchvision
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def get_video_info(video_path):
|
||||
"""Get video information using torchvision."""
|
||||
# Read video tensor (T, C, H, W)
|
||||
video_tensor, _, info = torchvision.io.read_video(str(video_path),
|
||||
output_format="TCHW",
|
||||
pts_unit="sec")
|
||||
|
||||
num_frames = video_tensor.shape[0]
|
||||
height = video_tensor.shape[2]
|
||||
width = video_tensor.shape[3]
|
||||
fps = info.get("video_fps", 0)
|
||||
duration = num_frames / fps if fps > 0 else 0
|
||||
|
||||
# Extract name
|
||||
_, _, videos_dir, video_name = str(video_path).split("/")
|
||||
|
||||
return {
|
||||
"path": str(video_name),
|
||||
"resolution": {
|
||||
"width": width,
|
||||
"height": height
|
||||
},
|
||||
"size": os.path.getsize(video_path),
|
||||
"fps": fps,
|
||||
"duration": duration,
|
||||
"num_frames": num_frames
|
||||
}
|
||||
|
||||
|
||||
def prepare_dataset_json(folder_path,
|
||||
output_name="videos2caption.json",
|
||||
num_workers=None) -> None:
|
||||
"""Prepare dataset information from a folder containing videos and prompt.txt."""
|
||||
folder_path = Path(folder_path)
|
||||
|
||||
# Read prompt file
|
||||
prompt_file = folder_path / "prompt.txt"
|
||||
if not prompt_file.exists():
|
||||
raise FileNotFoundError(f"prompt.txt not found in {folder_path}")
|
||||
|
||||
with open(prompt_file) as f:
|
||||
prompts = [line.strip() for line in f.readlines() if line.strip()]
|
||||
|
||||
# Read videos file
|
||||
videos_file = folder_path / "videos.txt"
|
||||
if not videos_file.exists():
|
||||
raise FileNotFoundError(f"videos.txt not found in {folder_path}")
|
||||
|
||||
with open(videos_file) as f:
|
||||
video_paths = [line.strip() for line in f.readlines() if line.strip()]
|
||||
|
||||
if len(prompts) != len(video_paths):
|
||||
raise ValueError(
|
||||
f"Number of prompts ({len(prompts)}) does not match number of videos ({len(video_paths)})"
|
||||
)
|
||||
|
||||
# Prepare arguments for multiprocessing
|
||||
process_args = [folder_path / video_path for video_path in video_paths]
|
||||
|
||||
# Determine number of workers
|
||||
if num_workers is None:
|
||||
num_workers = max(1, cpu_count() - 1) # Leave one CPU free
|
||||
|
||||
# Process videos in parallel
|
||||
start_time = time.time()
|
||||
with Pool(num_workers) as pool:
|
||||
results = list(
|
||||
tqdm(pool.imap(get_video_info, process_args),
|
||||
total=len(process_args),
|
||||
desc="Processing videos",
|
||||
unit="video"))
|
||||
|
||||
# Combine results with prompts
|
||||
dataset_info = []
|
||||
for result, prompt in zip(results, prompts):
|
||||
result["cap"] = [prompt]
|
||||
dataset_info.append(result)
|
||||
|
||||
# Calculate total processing time
|
||||
total_time = time.time() - start_time
|
||||
total_videos = len(dataset_info)
|
||||
avg_time_per_video = total_time / total_videos if total_videos > 0 else 0
|
||||
|
||||
print("\nProcessing completed:")
|
||||
print(f"Total videos processed: {total_videos}")
|
||||
print(f"Total time: {total_time:.2f} seconds")
|
||||
print(f"Average time per video: {avg_time_per_video:.2f} seconds")
|
||||
|
||||
# Save to JSON file
|
||||
output_file = folder_path / output_name
|
||||
with open(output_file, 'w') as f:
|
||||
json.dump(dataset_info, f, indent=2)
|
||||
|
||||
# Create merge.txt
|
||||
merge_file = folder_path / "merge.txt"
|
||||
with open(merge_file, 'w') as f:
|
||||
f.write(f"{folder_path}/videos,{output_file}\n")
|
||||
|
||||
print(f"Dataset information saved to {output_file}")
|
||||
print(f"Merge file created at {merge_file}")
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description='Prepare video dataset information in JSON format')
|
||||
parser.add_argument(
|
||||
'--folder',
|
||||
type=str,
|
||||
required=True,
|
||||
help='Path to the folder containing videos and prompt.txt')
|
||||
parser.add_argument(
|
||||
'--output',
|
||||
type=str,
|
||||
default='videos2caption.json',
|
||||
help='Name of the output JSON file (default: videos2caption.json)')
|
||||
parser.add_argument('--workers',
|
||||
type=int,
|
||||
default=32,
|
||||
help='Number of worker processes (default: 16)')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = parse_args()
|
||||
prepare_dataset_json(args.folder, args.output, args.workers)
|
||||
@@ -32,6 +32,7 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
|
||||
sp_world_size: int,
|
||||
global_rank: int,
|
||||
drop_last: bool = True,
|
||||
drop_first_row: bool = False,
|
||||
seed: int = 0,
|
||||
):
|
||||
self.batch_size = batch_size
|
||||
@@ -47,6 +48,11 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
|
||||
# Create a random permutation of all indices
|
||||
global_indices = torch.randperm(self.dataset_size, generator=rng)
|
||||
|
||||
if drop_first_row:
|
||||
# drop 0 in global_indices
|
||||
global_indices = global_indices[global_indices != 0]
|
||||
self.dataset_size = self.dataset_size - 1
|
||||
|
||||
if self.drop_last:
|
||||
# For drop_last=True, we:
|
||||
# 1. Ensure total samples is divisible by (batch_size * num_sp_groups)
|
||||
@@ -188,6 +194,7 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
cfg_rate: float = 0.0,
|
||||
seed: int = 42,
|
||||
drop_last: bool = True,
|
||||
drop_first_row: bool = False,
|
||||
text_padding_length: int = 512,
|
||||
):
|
||||
super().__init__()
|
||||
@@ -218,6 +225,7 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
sp_world_size=get_sp_world_size(),
|
||||
global_rank=get_world_rank(),
|
||||
drop_last=drop_last,
|
||||
drop_first_row=drop_first_row,
|
||||
seed=seed,
|
||||
)
|
||||
logger.info("Dataset initialized with %d parquet files and %d rows",
|
||||
@@ -280,6 +288,7 @@ def build_parquet_map_style_dataloader(
|
||||
num_data_workers,
|
||||
cfg_rate=0.0,
|
||||
drop_last=True,
|
||||
drop_first_row=False,
|
||||
text_padding_length=512,
|
||||
seed=42) -> Tuple[LatentsParquetMapStyleDataset, StatefulDataLoader]:
|
||||
dataset = LatentsParquetMapStyleDataset(
|
||||
@@ -287,6 +296,7 @@ def build_parquet_map_style_dataloader(
|
||||
batch_size,
|
||||
cfg_rate=cfg_rate,
|
||||
drop_last=drop_last,
|
||||
drop_first_row=drop_first_row,
|
||||
text_padding_length=text_padding_length,
|
||||
seed=seed)
|
||||
|
||||
|
||||
@@ -0,0 +1,592 @@
|
||||
# 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 DatasetBatch:
|
||||
"""
|
||||
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: DatasetBatch, **kwargs) -> DatasetBatch:
|
||||
"""
|
||||
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: DatasetBatch, **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: DatasetBatch, **kwargs) -> DatasetBatch:
|
||||
"""
|
||||
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: DatasetBatch, **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: DatasetBatch, **kwargs) -> DatasetBatch:
|
||||
"""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: DatasetBatch, **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: DatasetBatch, **kwargs) -> DatasetBatch:
|
||||
"""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: DatasetBatch, **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: DatasetBatch,
|
||||
temporal_sample_fn=None,
|
||||
**kwargs) -> DatasetBatch:
|
||||
"""
|
||||
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: DatasetBatch, **kwargs) -> DatasetBatch:
|
||||
"""
|
||||
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: DatasetBatch, **kwargs) -> DatasetBatch:
|
||||
"""
|
||||
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: DatasetBatch, **kwargs) -> DatasetBatch:
|
||||
"""
|
||||
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.
|
||||
|
||||
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.cfg)
|
||||
|
||||
def _load_raw_data(self) -> List[Dict]:
|
||||
"""Load raw data from JSON files."""
|
||||
all_data = []
|
||||
|
||||
# Read folder-annotation pairs
|
||||
with open(self.data_merge_path) as f:
|
||||
folder_anno_pairs = [
|
||||
line.strip().split(",") for line in f if line.strip()
|
||||
]
|
||||
|
||||
# Process each folder-annotation pair
|
||||
for folder, annotation_file in folder_anno_pairs:
|
||||
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"])
|
||||
|
||||
all_data.extend(data_items)
|
||||
|
||||
return all_data[self.start_idx:]
|
||||
|
||||
def _process_metadata(self) -> List[DatasetBatch]:
|
||||
"""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 = DatasetBatch(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: DatasetBatch,
|
||||
filter_counts: Dict[str, int]) -> bool:
|
||||
"""Apply all filter stages and update counters. Returns True if batch should be kept."""
|
||||
if not self.validation_stage.should_keep(batch):
|
||||
filter_counts["validation_failed"] += 1
|
||||
return False
|
||||
|
||||
if not self.resolution_filter_stage.should_keep(batch):
|
||||
filter_counts["resolution_failed"] += 1
|
||||
return False
|
||||
|
||||
if not self.frame_sampling_stage.should_keep(batch):
|
||||
filter_counts["frame_sampling_failed"] += 1
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def _log_filtering_stats(self, filter_counts: Dict[str, int],
|
||||
sample_num_frames: List[int], before_count: int,
|
||||
after_count: int):
|
||||
"""Log filtering statistics."""
|
||||
logger.info(
|
||||
"validation_failed: %d, resolution_failed: %d, frame_sampling_failed: %d, "
|
||||
"Counter(sample_num_frames): %s, before filter: %d, after filter: %d",
|
||||
filter_counts['validation_failed'],
|
||||
filter_counts['resolution_failed'],
|
||||
filter_counts['frame_sampling_failed'], Counter(sample_num_frames),
|
||||
before_count, after_count)
|
||||
|
||||
def __iter__(self):
|
||||
"""Iterate through processed data items."""
|
||||
for idx in range(len(self.processed_batches)):
|
||||
yield self._get_item(idx)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.processed_batches)
|
||||
|
||||
def _get_item(self, idx: int) -> Dict:
|
||||
"""Get a single processed data item."""
|
||||
batch = self.processed_batches[idx]
|
||||
|
||||
# Apply transformation stages
|
||||
batch = self.video_transform_stage.process(batch)
|
||||
batch = self.image_transform_stage.process(batch)
|
||||
batch = self.text_encoding_stage.process(batch)
|
||||
|
||||
# Build result dictionary
|
||||
result = {
|
||||
"pixel_values": batch.pixel_values,
|
||||
"text": batch.text,
|
||||
"input_ids": batch.input_ids,
|
||||
"cond_mask": batch.cond_mask,
|
||||
"path": batch.path,
|
||||
}
|
||||
|
||||
# Add video-specific fields
|
||||
if batch.is_video:
|
||||
result.update({"fps": batch.fps, "duration": batch.duration})
|
||||
|
||||
return result
|
||||
|
||||
def state_dict(self) -> Dict[str, Any]:
|
||||
"""Return state dict for checkpointing."""
|
||||
return {"processed_batches": self.processed_batches}
|
||||
|
||||
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
|
||||
"""Load state dict from checkpoint."""
|
||||
self.processed_batches = state_dict["processed_batches"]
|
||||
@@ -1,352 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from collections import Counter
|
||||
from os.path import join as opj
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from fastvideo.utils.dataset_utils import DecordInit
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
|
||||
|
||||
class SingletonMeta(type):
|
||||
_instances: dict[type, 'SingletonMeta'] = {}
|
||||
|
||||
def __call__(cls, *args, **kwargs):
|
||||
if cls not in cls._instances:
|
||||
instance = super().__call__(*args, **kwargs)
|
||||
cls._instances[cls] = instance
|
||||
return cls._instances[cls]
|
||||
|
||||
|
||||
class DataSetProg(metaclass=SingletonMeta):
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.cap_list: list[dict] = []
|
||||
self.elements: list[int] = []
|
||||
self.num_workers = 1
|
||||
self.n_elements = 0
|
||||
self.worker_elements: dict[int, list[int]] = {}
|
||||
self.n_used_elements: dict[int, int] = {}
|
||||
|
||||
def set_cap_list(self, num_workers, cap_list, n_elements) -> None:
|
||||
self.num_workers = num_workers
|
||||
self.cap_list = cap_list
|
||||
self.n_elements = n_elements
|
||||
self.elements = list(range(n_elements))
|
||||
random.shuffle(self.elements)
|
||||
print(f"n_elements: {len(self.elements)}", flush=True)
|
||||
|
||||
for i in range(self.num_workers):
|
||||
self.n_used_elements[i] = 0
|
||||
per_worker = int(
|
||||
math.ceil(len(self.elements) / float(self.num_workers)))
|
||||
start = i * per_worker
|
||||
end = min(start + per_worker, len(self.elements))
|
||||
self.worker_elements[i] = self.elements[start:end]
|
||||
|
||||
def get_item(self, work_info) -> int:
|
||||
worker_id = 0 if work_info is None else work_info.id
|
||||
|
||||
idx = self.worker_elements[worker_id][
|
||||
self.n_used_elements[worker_id] %
|
||||
len(self.worker_elements[worker_id])]
|
||||
self.n_used_elements[worker_id] += 1
|
||||
return idx
|
||||
|
||||
|
||||
dataset_prog = DataSetProg()
|
||||
|
||||
|
||||
def filter_resolution(h: int,
|
||||
w: int,
|
||||
max_h_div_w_ratio: float = 17 / 16,
|
||||
min_h_div_w_ratio: float = 8 / 16) -> bool:
|
||||
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
|
||||
|
||||
|
||||
class T2V_dataset(Dataset):
|
||||
|
||||
def __init__(self,
|
||||
args,
|
||||
transform,
|
||||
temporal_sample,
|
||||
tokenizer,
|
||||
transform_topcrop,
|
||||
start_idx=0) -> None:
|
||||
self.start_idx = start_idx
|
||||
self.data = args.data_merge_path
|
||||
self.num_frames = args.num_frames
|
||||
self.train_fps = args.train_fps
|
||||
self.use_image_num = args.use_image_num
|
||||
self.transform = transform
|
||||
self.transform_topcrop = transform_topcrop
|
||||
self.temporal_sample = temporal_sample
|
||||
self.tokenizer = tokenizer
|
||||
self.text_max_length = args.text_max_length
|
||||
self.cfg = args.cfg
|
||||
self.speed_factor = args.speed_factor
|
||||
self.max_height = args.max_height
|
||||
self.max_width = args.max_width
|
||||
self.drop_short_ratio = args.drop_short_ratio
|
||||
assert self.speed_factor >= 1
|
||||
self.v_decoder = DecordInit()
|
||||
self.video_length_tolerance_range = args.video_length_tolerance_range
|
||||
self.support_Chinese = True
|
||||
if "mt5" not in args.text_encoder_name:
|
||||
self.support_Chinese = False
|
||||
|
||||
cap_list = self.get_cap_list()
|
||||
|
||||
assert len(cap_list) > 0
|
||||
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
|
||||
self.lengths = self.sample_num_frames
|
||||
|
||||
n_elements = len(cap_list)
|
||||
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list,
|
||||
n_elements)
|
||||
|
||||
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
|
||||
|
||||
def set_checkpoint(self, n_used_elements):
|
||||
for i in range(len(dataset_prog.n_used_elements)):
|
||||
dataset_prog.n_used_elements[i] = n_used_elements
|
||||
|
||||
def __len__(self):
|
||||
return dataset_prog.n_elements
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
data = self.get_data(idx)
|
||||
return data
|
||||
|
||||
def get_data(self, idx) -> dict:
|
||||
path = dataset_prog.cap_list[idx]["path"]
|
||||
if path.endswith(".mp4"):
|
||||
return self.get_video(idx)
|
||||
else:
|
||||
return self.get_image(idx)
|
||||
|
||||
def get_video(self, idx) -> dict:
|
||||
video_path = dataset_prog.cap_list[idx]["path"]
|
||||
assert os.path.exists(video_path), f"file {video_path} do not exist!"
|
||||
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
|
||||
|
||||
torchvision_video, _, metadata = torchvision.io.read_video(
|
||||
video_path, output_format="TCHW")
|
||||
video = torchvision_video[frame_indices]
|
||||
video = self.transform(video)
|
||||
video = rearrange(video, "t c h w -> c t h w")
|
||||
video = video.to(torch.uint8)
|
||||
assert video.dtype == torch.uint8
|
||||
|
||||
h, w = video.shape[-2:]
|
||||
assert (
|
||||
h / w <= 17 / 16 and h / w >= 8 / 16
|
||||
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
|
||||
|
||||
video = video.float() / 127.5 - 1.0
|
||||
|
||||
text = dataset_prog.cap_list[idx]["cap"]
|
||||
if not isinstance(text, list):
|
||||
text = [text]
|
||||
text = [random.choice(text)]
|
||||
|
||||
text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"]
|
||||
cond_mask = text_tokens_and_mask["attention_mask"]
|
||||
return dict(pixel_values=video,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=video_path,
|
||||
fps=dataset_prog.cap_list[idx]["fps"],
|
||||
duration=dataset_prog.cap_list[idx]["duration"])
|
||||
|
||||
def get_image(self, idx) -> dict:
|
||||
image_data = dataset_prog.cap_list[
|
||||
idx] # [{'path': path, 'cap': cap}, ...]
|
||||
|
||||
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
|
||||
image = torch.from_numpy(np.array(image)) # [h, w, c]
|
||||
image = rearrange(image, "h w c -> c h w").unsqueeze(0) # [1 c h w]
|
||||
# for i in image:
|
||||
# h, w = i.shape[-2:]
|
||||
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
|
||||
|
||||
image = (self.transform_topcrop(image) if "human_images"
|
||||
in image_data["path"] else self.transform(image)
|
||||
) # [1 C H W] -> num_img [1 C H W]
|
||||
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
|
||||
|
||||
image = image.float() / 127.5 - 1.0
|
||||
|
||||
caps: list[str] = (image_data["cap"] if isinstance(
|
||||
image_data["cap"], list) else [image_data["cap"]])
|
||||
caps = [random.choice(caps)]
|
||||
text = caps
|
||||
input_ids, cond_mask = [], []
|
||||
single_text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
single_text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"] # 1, l
|
||||
cond_mask = text_tokens_and_mask["attention_mask"] # 1, l
|
||||
return dict(
|
||||
pixel_values=image,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=image_data["path"],
|
||||
)
|
||||
|
||||
def define_frame_index(self, cap_list) -> tuple[list[dict], list[int]]:
|
||||
new_cap_list = []
|
||||
sample_num_frames = []
|
||||
cnt_too_long = 0
|
||||
cnt_too_short = 0
|
||||
cnt_no_cap = 0
|
||||
cnt_no_resolution = 0
|
||||
cnt_resolution_mismatch = 0
|
||||
cnt_movie = 0
|
||||
cnt_img = 0
|
||||
for i in cap_list:
|
||||
path = i["path"]
|
||||
cap = i.get("cap", None)
|
||||
# ======no caption=====
|
||||
if cap is None:
|
||||
cnt_no_cap += 1
|
||||
continue
|
||||
if path.endswith(".mp4"):
|
||||
# ======no fps and duration=====
|
||||
duration = i.get("duration", None)
|
||||
fps = i.get("fps", None)
|
||||
if fps is None or duration is None:
|
||||
continue
|
||||
|
||||
# ======resolution mismatch=====
|
||||
resolution = i.get("resolution", None)
|
||||
if resolution is None:
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
else:
|
||||
if (resolution.get("height", None) is None
|
||||
or resolution.get("width", None) is None):
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
height, width = i["resolution"]["height"], i["resolution"][
|
||||
"width"]
|
||||
aspect = self.max_height / self.max_width
|
||||
hw_aspect_thr = 1.5
|
||||
is_pick = filter_resolution(
|
||||
height,
|
||||
width,
|
||||
max_h_div_w_ratio=hw_aspect_thr * aspect,
|
||||
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
|
||||
)
|
||||
if not is_pick:
|
||||
print("resolution mismatch")
|
||||
cnt_resolution_mismatch += 1
|
||||
continue
|
||||
|
||||
# if path == 'finetrainers/3dgs-dissolve/videos/1.mp4':
|
||||
# from IPython import embed; embed()
|
||||
i["num_frames"] = math.ceil(fps * duration)
|
||||
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
|
||||
if i["num_frames"] / fps > self.video_length_tolerance_range * (
|
||||
self.num_frames / self.train_fps * self.speed_factor
|
||||
): # too long video is not suitable for this training stage (self.num_frames)
|
||||
cnt_too_long += 1
|
||||
continue
|
||||
|
||||
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
|
||||
frame_interval = fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, i["num_frames"],
|
||||
frame_interval).astype(int)
|
||||
|
||||
# comment out it to enable dynamic frames training
|
||||
if (len(frame_indices) < self.num_frames
|
||||
and random.random() < self.drop_short_ratio):
|
||||
cnt_too_short += 1
|
||||
continue
|
||||
|
||||
# too long video will be temporal-crop randomly
|
||||
if len(frame_indices) > self.num_frames:
|
||||
begin_index, end_index = self.temporal_sample(
|
||||
len(frame_indices))
|
||||
frame_indices = frame_indices[begin_index:end_index]
|
||||
# frame_indices = frame_indices[:self.num_frames] # head crop
|
||||
i["sample_frame_index"] = frame_indices.tolist()
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = len(
|
||||
i["sample_frame_index"]
|
||||
) # will use in dataloader(group sampler)
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
elif path.endswith(".jpg"): # image
|
||||
cnt_img += 1
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = 1
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
else:
|
||||
raise NameError(
|
||||
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
|
||||
)
|
||||
# import ipdb;ipdb.set_trace()
|
||||
main_print(
|
||||
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
|
||||
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
|
||||
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
|
||||
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}"
|
||||
)
|
||||
return new_cap_list, sample_num_frames
|
||||
|
||||
def decord_read(self, path, frame_indices) -> torch.Tensor:
|
||||
decord_vr = self.v_decoder(path)
|
||||
video_data = decord_vr.get_batch(frame_indices).asnumpy()
|
||||
video_data = torch.from_numpy(video_data)
|
||||
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
|
||||
return video_data
|
||||
|
||||
def read_jsons(self, data) -> list[dict]:
|
||||
cap_lists = []
|
||||
with open(data) as f:
|
||||
folder_anno = [
|
||||
i.strip().split(",") for i in f.readlines()
|
||||
if len(i.strip()) > 0
|
||||
]
|
||||
print(folder_anno)
|
||||
for folder, anno in folder_anno:
|
||||
with open(anno) as f:
|
||||
sub_list = json.load(f)
|
||||
for i in range(len(sub_list)):
|
||||
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
|
||||
cap_lists += sub_list
|
||||
return cap_lists
|
||||
|
||||
def get_cap_list(self) -> list:
|
||||
cap_lists = self.read_jsons(self.data)[self.start_idx:]
|
||||
return cap_lists
|
||||
@@ -0,0 +1,100 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from: https://github.com/a-r-r-o-w/finetrainers/blob/main/finetrainers/data/dataset.py
|
||||
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)
|
||||
|
||||
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"]
|
||||
if not pathlib.Path(image_path).is_file(
|
||||
) and not image_path.startswith("http"):
|
||||
logger.warning("Image file %s does not exist.",
|
||||
image_path.as_posix())
|
||||
else:
|
||||
sample["image"] = load_image(sample["image_path"])
|
||||
|
||||
if sample.get("video_path", None) is not None:
|
||||
video_path = sample["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.as_posix())
|
||||
else:
|
||||
sample["video"] = load_video(sample["video_path"])
|
||||
|
||||
if sample.get("control_image_path", None) is not None:
|
||||
control_image_path = sample["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.as_posix())
|
||||
else:
|
||||
sample["control_image"] = load_image(
|
||||
sample["control_image_path"])
|
||||
|
||||
if sample.get("control_video_path", None) is not None:
|
||||
control_video_path = sample["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(
|
||||
sample["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)
|
||||
|
||||
@@ -379,7 +388,8 @@ class TrainingArgs(FastVideoArgs):
|
||||
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
|
||||
@@ -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")
|
||||
|
||||
@@ -11,10 +11,10 @@ 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
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.v1.layers.attention import AttentionMetadata
|
||||
from fastvideo.v1.attention import AttentionMetadata
|
||||
from fastvideo.v1.pipelines import ForwardBatch
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -37,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. "
|
||||
@@ -55,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.
|
||||
|
||||
@@ -1,19 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.v1.layers.attention.backends.abstract import (
|
||||
AttentionBackend, AttentionMetadata, AttentionMetadataBuilder)
|
||||
from fastvideo.v1.layers.attention.layer import (DistributedAttention,
|
||||
DistributedAttention_VSA,
|
||||
LocalAttention)
|
||||
from fastvideo.v1.layers.attention.selector import get_attn_backend
|
||||
|
||||
__all__ = [
|
||||
"DistributedAttention",
|
||||
"LocalAttention",
|
||||
"DistributedAttention_VSA",
|
||||
"AttentionBackend",
|
||||
"AttentionMetadata",
|
||||
"AttentionMetadataBuilder",
|
||||
# "AttentionState",
|
||||
"get_attn_backend",
|
||||
]
|
||||
@@ -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
|
||||
|
||||
@@ -95,6 +96,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 +129,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 +139,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 +186,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 +212,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
|
||||
|
||||
@@ -6,7 +6,7 @@ import torch
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
# TODO
|
||||
@@ -14,12 +14,13 @@ class BaseDiT(nn.Module, ABC):
|
||||
_fsdp_shard_conditions: list = []
|
||||
_compile_conditions: list = []
|
||||
_param_names_mapping: dict
|
||||
_reverse_param_names_mapping: dict
|
||||
hidden_size: int
|
||||
num_attention_heads: int
|
||||
num_channels_latents: int
|
||||
# always supports torch_sdpa
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = DiTConfig()._supported_attention_backends
|
||||
AttentionBackendEnum, ...] = DiTConfig()._supported_attention_backends
|
||||
|
||||
def __init_subclass__(cls) -> None:
|
||||
required_class_attrs = [
|
||||
@@ -65,7 +66,7 @@ class BaseDiT(nn.Module, ABC):
|
||||
)
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
|
||||
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
|
||||
@@ -78,6 +79,7 @@ class CachableDiT(BaseDiT):
|
||||
# These are required class attributes that should be overridden by concrete implementations
|
||||
_fsdp_shard_conditions = []
|
||||
_param_names_mapping = {}
|
||||
_reverse_param_names_mapping = {}
|
||||
_lora_param_names_mapping: dict = {}
|
||||
# Ensure these instance attributes are properly defined in subclasses
|
||||
hidden_size: int
|
||||
@@ -85,7 +87,7 @@ class CachableDiT(BaseDiT):
|
||||
num_channels_latents: int
|
||||
# always supports torch_sdpa
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = DiTConfig()._supported_attention_backends
|
||||
AttentionBackendEnum, ...] = DiTConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, config: DiTConfig, **kwargs) -> None:
|
||||
super().__init__(config, **kwargs)
|
||||
|
||||
@@ -6,11 +6,11 @@ import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.v1.attention import DistributedAttention, LocalAttention
|
||||
from fastvideo.v1.configs.models.dits import HunyuanVideoConfig
|
||||
from fastvideo.v1.configs.sample.teacache import TeaCacheParams
|
||||
from fastvideo.v1.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.v1.forward_context import get_forward_context
|
||||
from fastvideo.v1.layers.attention import DistributedAttention, LocalAttention
|
||||
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
|
||||
ScaleResidualLayerNormScaleShift)
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
@@ -23,7 +23,7 @@ from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
|
||||
unpatchify)
|
||||
from fastvideo.v1.models.dits.base import CachableDiT
|
||||
from fastvideo.v1.models.utils import modulate
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class HunyuanRMSNorm(nn.Module):
|
||||
@@ -96,7 +96,8 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
@@ -303,7 +304,8 @@ class MMSingleStreamBlock(nn.Module):
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
@@ -440,6 +442,8 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
|
||||
_supported_attention_backends = HunyuanVideoConfig(
|
||||
)._supported_attention_backends
|
||||
_param_names_mapping = HunyuanVideoConfig()._param_names_mapping
|
||||
_reverse_param_names_mapping = HunyuanVideoConfig(
|
||||
)._reverse_param_names_mapping
|
||||
_lora_param_names_mapping = HunyuanVideoConfig()._lora_param_names_mapping
|
||||
|
||||
def __init__(self, config: HunyuanVideoConfig, hf_config: dict[str, Any]):
|
||||
@@ -876,8 +880,8 @@ class IndividualTokenRefinerBlock(nn.Module):
|
||||
num_heads=num_attention_heads,
|
||||
head_size=hidden_size // num_attention_heads,
|
||||
# TODO: remove hardcode; remove STA
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA),
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA),
|
||||
)
|
||||
|
||||
def forward(self, x, c):
|
||||
|
||||
@@ -16,9 +16,9 @@ import torch
|
||||
from einops import rearrange, repeat
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.v1.attention import DistributedAttention, LocalAttention
|
||||
from fastvideo.v1.configs.models.dits import StepVideoConfig
|
||||
from fastvideo.v1.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.v1.layers.attention import DistributedAttention, LocalAttention
|
||||
from fastvideo.v1.layers.layernorm import LayerNormScaleShift
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
from fastvideo.v1.layers.mlp import MLP
|
||||
@@ -26,7 +26,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
|
||||
get_rotary_pos_embed)
|
||||
from fastvideo.v1.layers.visual_embedding import TimestepEmbedder
|
||||
from fastvideo.v1.models.dits.base import BaseDiT
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class PatchEmbed2D(nn.Module):
|
||||
@@ -139,16 +139,17 @@ class StepVideoRMSNorm(nn.Module):
|
||||
|
||||
class SelfAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
hidden_dim,
|
||||
head_dim,
|
||||
rope_split: Tuple[int, int, int] = (64, 32, 32),
|
||||
bias: bool = False,
|
||||
with_rope: bool = True,
|
||||
with_qk_norm: bool = True,
|
||||
attn_type: str = "torch",
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA)):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_dim,
|
||||
head_dim,
|
||||
rope_split: Tuple[int, int, int] = (64, 32, 32),
|
||||
bias: bool = False,
|
||||
with_rope: bool = True,
|
||||
with_qk_norm: bool = True,
|
||||
attn_type: str = "torch",
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA)):
|
||||
super().__init__()
|
||||
self.head_dim = head_dim
|
||||
self.hidden_dim = hidden_dim
|
||||
@@ -257,7 +258,8 @@ class CrossAttention(nn.Module):
|
||||
head_dim,
|
||||
bias=False,
|
||||
with_qk_norm=True,
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA)
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.head_dim = head_dim
|
||||
@@ -458,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
|
||||
|
||||
@@ -8,15 +8,14 @@ import torch
|
||||
import torch.nn as nn
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.attention import (DistributedAttention,
|
||||
DistributedAttention_VSA, LocalAttention)
|
||||
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.attention import (DistributedAttention,
|
||||
DistributedAttention_VSA,
|
||||
LocalAttention)
|
||||
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
|
||||
@@ -27,7 +26,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
|
||||
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
|
||||
PatchEmbed, TimestepEmbedder)
|
||||
from fastvideo.v1.models.dits.base import CachableDiT
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class WanImageEmbedding(torch.nn.Module):
|
||||
@@ -35,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:
|
||||
@@ -126,8 +125,8 @@ class WanSelfAttention(nn.Module):
|
||||
dropout_rate=0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA))
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA))
|
||||
|
||||
def forward(self, x: torch.Tensor, context: torch.Tensor,
|
||||
context_lens: int):
|
||||
@@ -175,7 +174,8 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
eps=1e-6,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None
|
||||
) -> None:
|
||||
super().__init__(dim, num_heads, window_size, qk_norm, eps,
|
||||
supported_attention_backends)
|
||||
@@ -217,21 +217,22 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
|
||||
class WanTransformerBlock(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
prefix: str = ""):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -262,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:
|
||||
@@ -282,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")
|
||||
@@ -359,21 +362,22 @@ class WanTransformerBlock(nn.Module):
|
||||
|
||||
class WanTransformerBlock_VSA(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
prefix: str = ""):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -405,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:
|
||||
@@ -425,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")
|
||||
@@ -512,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,
|
||||
@@ -562,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(
|
||||
@@ -570,7 +578,7 @@ class WanTransformer3DModel(CachableDiT):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
# Initialize cache-related attributes
|
||||
# For type checking
|
||||
self.previous_e0_even = None
|
||||
self.previous_e0_odd = None
|
||||
self.previous_residual_even = None
|
||||
@@ -581,7 +589,6 @@ class WanTransformer3DModel(CachableDiT):
|
||||
self.accumulated_rel_l1_distance_even = 0
|
||||
self.accumulated_rel_l1_distance_odd = 0
|
||||
self.cnt = 0
|
||||
|
||||
self.__post_init__()
|
||||
|
||||
def forward(self,
|
||||
@@ -668,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,
|
||||
|
||||
@@ -8,12 +8,13 @@ from torch import nn
|
||||
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
|
||||
ImageEncoderConfig,
|
||||
TextEncoderConfig)
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class TextEncoder(nn.Module, ABC):
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = TextEncoderConfig()._supported_attention_backends
|
||||
AttentionBackendEnum,
|
||||
...] = TextEncoderConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, config: TextEncoderConfig) -> None:
|
||||
super().__init__()
|
||||
@@ -34,13 +35,14 @@ class TextEncoder(nn.Module, ABC):
|
||||
pass
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
|
||||
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
|
||||
class ImageEncoder(nn.Module, ABC):
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = ImageEncoderConfig()._supported_attention_backends
|
||||
AttentionBackendEnum,
|
||||
...] = ImageEncoderConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, config: ImageEncoderConfig) -> None:
|
||||
super().__init__()
|
||||
@@ -56,5 +58,5 @@ class ImageEncoder(nn.Module, ABC):
|
||||
pass
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
|
||||
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
@@ -8,13 +8,13 @@ from typing import Iterable, Optional, Set, Tuple, Union
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
# from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask, _prepare_4d_attention_mask
|
||||
from fastvideo.v1.attention import LocalAttention
|
||||
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
|
||||
CLIPTextConfig,
|
||||
CLIPVisionConfig)
|
||||
from fastvideo.v1.distributed import divide, get_tp_world_size
|
||||
from fastvideo.v1.layers.activation import get_act_fn
|
||||
# from transformers.modeling_attn_mask_utils import _create_4d_causal_attention_mask, _prepare_4d_attention_mask
|
||||
from fastvideo.v1.layers.attention import LocalAttention
|
||||
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
|
||||
RowParallelLinear)
|
||||
from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
|
||||
@@ -28,12 +28,12 @@ from typing import Any, Dict, Iterable, Optional, Set, Tuple
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
# from vllm.model_executor.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.attention import LocalAttention
|
||||
# from ..utils import (extract_layer_index)
|
||||
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput, LlamaConfig
|
||||
from fastvideo.v1.distributed import get_tp_world_size
|
||||
from fastvideo.v1.layers.activation import SiluAndMul
|
||||
# from vllm.model_executor.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.layers.attention import LocalAttention
|
||||
from fastvideo.v1.layers.layernorm import RMSNorm
|
||||
from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
|
||||
QKVParallelLinear, RowParallelLinear)
|
||||
|
||||
@@ -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,7 @@ 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
|
||||
|
||||
# Text inputs
|
||||
prompt: Optional[Union[str, List[str]]] = None
|
||||
@@ -136,3 +139,30 @@ 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
|
||||
info: Optional[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
|
||||
|
||||
# Training loss
|
||||
loss: torch.Tensor | None = None
|
||||
|
||||
# Training outputs
|
||||
total_loss: float | None = None
|
||||
grad_norm: float | None = None
|
||||
|
||||
@@ -3,6 +3,7 @@ import gc
|
||||
import multiprocessing
|
||||
import os
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
from itertools import chain
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import numpy as np
|
||||
@@ -10,17 +11,16 @@ 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.distributed import get_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages import TextEncodingStage
|
||||
from fastvideo.v1.pipelines.stages import EncodingStage, TextEncodingStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -36,6 +36,9 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="image_encoding_stage",
|
||||
stage=EncodingStage(vae=self.get_module("vae"), ))
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
@@ -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],
|
||||
@@ -103,7 +106,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 +116,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,
|
||||
)
|
||||
@@ -285,7 +283,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 +294,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",
|
||||
|
||||
@@ -19,7 +19,7 @@ logger = init_logger(__name__)
|
||||
def main(args) -> None:
|
||||
args.model_path = maybe_download_model(args.model_path)
|
||||
maybe_init_distributed_environment_and_model_parallel(1, 1)
|
||||
num_gpus = os.environ["WORLD_SIZE"]
|
||||
num_gpus = int(os.environ["WORLD_SIZE"])
|
||||
assert num_gpus == 1, "Only support 1 GPU"
|
||||
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
|
||||
kwargs = {
|
||||
@@ -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",
|
||||
|
||||
@@ -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,38 +3,42 @@
|
||||
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 (
|
||||
sequence_model_parallel_all_gather)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.layers.attention import get_attn_backend
|
||||
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.platforms import _Backend
|
||||
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
|
||||
|
||||
from fastvideo.v1.layers.attention.backends.sliding_tile_attn import (
|
||||
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
|
||||
from fastvideo.v1.layers.attention.backends.video_sparse_attn import (
|
||||
try:
|
||||
from fastvideo.v1.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionBackend)
|
||||
vsa_available = True
|
||||
except ImportError:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -55,10 +59,11 @@ class DenoisingStage(PipelineStage):
|
||||
self.attn_backend = get_attn_backend(
|
||||
head_size=attn_head_size,
|
||||
dtype=torch.float16, # TODO(will): hack
|
||||
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN,
|
||||
_Backend.VIDEO_SPARSE_ATTN,
|
||||
_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA) # hack
|
||||
supported_attention_backends=(
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA
|
||||
) # hack
|
||||
)
|
||||
|
||||
def forward(
|
||||
@@ -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__)
|
||||
@@ -58,7 +60,7 @@ class EncodingStage(PipelineStage):
|
||||
latent_height = batch.height // self.vae.spatial_compression_ratio
|
||||
latent_width = batch.width // self.vae.spatial_compression_ratio
|
||||
|
||||
image = load_image(image_path)
|
||||
image = batch.pil_image
|
||||
image = self.preprocess(
|
||||
image,
|
||||
vae_scale_factor=self.vae.spatial_compression_ratio,
|
||||
@@ -95,7 +97,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 +176,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, V.not_none)
|
||||
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
|
||||
|
||||
@@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Optional
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
# imported by other files, do not remove
|
||||
from fastvideo.v1.platforms.interface import _Backend # noqa: F401
|
||||
from fastvideo.v1.platforms.interface import AttentionBackendEnum # noqa: F401
|
||||
from fastvideo.v1.platforms.interface import Platform, PlatformEnum
|
||||
from fastvideo.v1.utils import resolve_obj_by_qualname
|
||||
|
||||
|
||||
@@ -13,8 +13,9 @@ from typing_extensions import ParamSpec
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.platforms.interface import (DeviceCapability, Platform,
|
||||
PlatformEnum, _Backend)
|
||||
from fastvideo.v1.platforms.interface import (AttentionBackendEnum,
|
||||
DeviceCapability, Platform,
|
||||
PlatformEnum)
|
||||
from fastvideo.v1.utils import import_pynvml
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -106,79 +107,85 @@ class CudaPlatformBase(Platform):
|
||||
return float(torch.cuda.max_memory_allocated(device))
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend: Optional[_Backend],
|
||||
def get_attn_backend_cls(cls,
|
||||
selected_backend: Optional[AttentionBackendEnum],
|
||||
head_size: int, dtype: torch.dtype) -> str:
|
||||
# TODO(will): maybe come up with a more general interface for local attention
|
||||
# if distributed is False, we always try to use Flash attn
|
||||
|
||||
logger.info("Trying FASTVIDEO_ATTENTION_BACKEND=%s",
|
||||
envs.FASTVIDEO_ATTENTION_BACKEND)
|
||||
if selected_backend == _Backend.SLIDING_TILE_ATTN:
|
||||
if selected_backend == AttentionBackendEnum.SLIDING_TILE_ATTN:
|
||||
try:
|
||||
from st_attn import sliding_tile_attention # noqa: F401
|
||||
|
||||
from fastvideo.v1.layers.attention.backends.sliding_tile_attn import ( # noqa: F401
|
||||
from fastvideo.v1.attention.backends.sliding_tile_attn import ( # noqa: F401
|
||||
SlidingTileAttentionBackend)
|
||||
logger.info("Using Sliding Tile Attention backend.")
|
||||
return "fastvideo.v1.layers.attention.backends.sliding_tile_attn.SlidingTileAttentionBackend"
|
||||
|
||||
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."
|
||||
)
|
||||
elif selected_backend == _Backend.SAGE_ATTN:
|
||||
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
|
||||
|
||||
from fastvideo.v1.layers.attention.backends.sage_attn import ( # noqa: F401
|
||||
from fastvideo.v1.attention.backends.sage_attn import ( # noqa: F401
|
||||
SageAttentionBackend)
|
||||
logger.info("Using Sage Attention backend.")
|
||||
return "fastvideo.v1.layers.attention.backends.sage_attn.SageAttentionBackend"
|
||||
|
||||
return "fastvideo.v1.attention.backends.sage_attn.SageAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
logger.info(
|
||||
"Sage Attention backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == _Backend.VIDEO_SPARSE_ATTN:
|
||||
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
|
||||
try:
|
||||
from vsa import block_sparse_attn # noqa: F401
|
||||
|
||||
from fastvideo.v1.layers.attention.backends.video_sparse_attn import ( # noqa: F401
|
||||
from fastvideo.v1.attention.backends.video_sparse_attn import ( # noqa: F401
|
||||
VideoSparseAttentionBackend)
|
||||
logger.info("Using Video Sparse Attention backend.")
|
||||
return "fastvideo.v1.layers.attention.backends.video_sparse_attn.VideoSparseAttentionBackend"
|
||||
|
||||
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."
|
||||
)
|
||||
elif selected_backend == _Backend.TORCH_SDPA:
|
||||
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.layers.attention.backends.sdpa.SDPABackend"
|
||||
elif selected_backend == _Backend.FLASH_ATTN or selected_backend is None:
|
||||
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
|
||||
elif selected_backend == AttentionBackendEnum.FLASH_ATTN or selected_backend is None:
|
||||
pass
|
||||
elif selected_backend:
|
||||
raise ValueError(f"Invalid attention backend for {cls.device_name}")
|
||||
|
||||
target_backend = _Backend.FLASH_ATTN
|
||||
target_backend = AttentionBackendEnum.FLASH_ATTN
|
||||
if not cls.has_device_capability(80):
|
||||
logger.info(
|
||||
"Cannot use FlashAttention-2 backend for Volta and Turing "
|
||||
"GPUs.")
|
||||
target_backend = _Backend.TORCH_SDPA
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
elif dtype not in (torch.float16, torch.bfloat16):
|
||||
logger.info(
|
||||
"Cannot use FlashAttention-2 backend for dtype other than "
|
||||
"torch.float16 or torch.bfloat16.")
|
||||
target_backend = _Backend.TORCH_SDPA
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
|
||||
# FlashAttn is valid for the model, checking if the package is
|
||||
# installed.
|
||||
if target_backend == _Backend.FLASH_ATTN:
|
||||
if target_backend == AttentionBackendEnum.FLASH_ATTN:
|
||||
try:
|
||||
import flash_attn # noqa: F401
|
||||
|
||||
from fastvideo.v1.layers.attention.backends.flash_attn import ( # noqa: F401
|
||||
from fastvideo.v1.attention.backends.flash_attn import ( # noqa: F401
|
||||
FlashAttentionBackend)
|
||||
|
||||
supported_sizes = \
|
||||
@@ -187,20 +194,22 @@ class CudaPlatformBase(Platform):
|
||||
logger.info(
|
||||
"Cannot use FlashAttention-2 backend for head size %d.",
|
||||
head_size)
|
||||
target_backend = _Backend.TORCH_SDPA
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
except ImportError:
|
||||
logger.info("Cannot use FlashAttention-2 backend because the "
|
||||
"flash_attn package is not found. "
|
||||
"Make sure that flash_attn was built and installed "
|
||||
"(on by default).")
|
||||
target_backend = _Backend.TORCH_SDPA
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
|
||||
if target_backend == _Backend.TORCH_SDPA:
|
||||
if target_backend == AttentionBackendEnum.TORCH_SDPA:
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
return "fastvideo.v1.layers.attention.backends.sdpa.SDPABackend"
|
||||
|
||||
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
|
||||
|
||||
logger.info("Using Flash Attention backend.")
|
||||
return "fastvideo.v1.layers.attention.backends.flash_attn.FlashAttentionBackend"
|
||||
|
||||
return "fastvideo.v1.attention.backends.flash_attn.FlashAttentionBackend"
|
||||
|
||||
@classmethod
|
||||
def get_device_communicator_cls(cls) -> str:
|
||||
|
||||
@@ -13,7 +13,7 @@ from fastvideo.v1.logger import init_logger
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class _Backend(enum.Enum):
|
||||
class AttentionBackendEnum(enum.Enum):
|
||||
FLASH_ATTN = enum.auto()
|
||||
SLIDING_TILE_ATTN = enum.auto()
|
||||
TORCH_SDPA = enum.auto()
|
||||
@@ -88,7 +88,8 @@ class Platform:
|
||||
return self._enum == PlatformEnum.CUDA
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend: Optional[_Backend],
|
||||
def get_attn_backend_cls(cls,
|
||||
selected_backend: Optional[AttentionBackendEnum],
|
||||
head_size: int, dtype: torch.dtype) -> str:
|
||||
"""Get the attention backend class of a device."""
|
||||
return ""
|
||||
@@ -168,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:
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
import os
|
||||
import sys
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
import pytest
|
||||
|
||||
NUM_NODES = "1"
|
||||
NUM_GPUS_PER_NODE = "1"
|
||||
|
||||
# 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", "1",
|
||||
"--tp-size", "1",
|
||||
"--num-gpus", "1",
|
||||
"--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,98 @@
|
||||
import modal
|
||||
|
||||
app = modal.App()
|
||||
|
||||
import os
|
||||
|
||||
image_version = os.getenv("IMAGE_VERSION", "latest")
|
||||
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")
|
||||
.apt_install("cmake", "pkg-config", "build-essential", "curl", "libssl-dev")
|
||||
.run_commands("curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y --default-toolchain stable")
|
||||
.run_commands("echo 'source ~/.cargo/env' >> ~/.bashrc")
|
||||
.env({"PATH": "/root/.cargo/bin:$PATH"})
|
||||
.run_commands("/bin/bash -c 'source $HOME/.local/bin/env && source /opt/venv/bin/activate && cd /FastVideo && uv pip install -e .[test]'")
|
||||
)
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1800)
|
||||
def run_encoder_tests():
|
||||
"""Run encoder tests on L40S GPU"""
|
||||
import subprocess
|
||||
import sys
|
||||
import os
|
||||
|
||||
os.chdir("/FastVideo")
|
||||
|
||||
command = """
|
||||
source /opt/venv/bin/activate &&
|
||||
pytest ./fastvideo/v1/tests/encoders -s
|
||||
"""
|
||||
|
||||
result = subprocess.run([
|
||||
"/bin/bash", "-c", command
|
||||
], stdout=sys.stdout, stderr=sys.stderr, check=False)
|
||||
|
||||
sys.exit(result.returncode)
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1800)
|
||||
def run_vae_tests():
|
||||
"""Run VAE tests on L40S GPU"""
|
||||
import subprocess
|
||||
import sys
|
||||
import os
|
||||
|
||||
os.chdir("/FastVideo")
|
||||
|
||||
command = """
|
||||
source /opt/venv/bin/activate &&
|
||||
pytest ./fastvideo/v1/tests/vaes -s
|
||||
"""
|
||||
|
||||
result = subprocess.run([
|
||||
"/bin/bash", "-c", command
|
||||
], stdout=sys.stdout, stderr=sys.stderr, check=False)
|
||||
|
||||
sys.exit(result.returncode)
|
||||
|
||||
@app.function(gpu="L40S:1", image=image, timeout=1800)
|
||||
def run_transformer_tests():
|
||||
"""Run transformer tests on L40S GPU"""
|
||||
import subprocess
|
||||
import sys
|
||||
import os
|
||||
|
||||
os.chdir("/FastVideo")
|
||||
|
||||
command = """
|
||||
source /opt/venv/bin/activate &&
|
||||
pytest ./fastvideo/v1/tests/transformers -s
|
||||
"""
|
||||
|
||||
result = subprocess.run([
|
||||
"/bin/bash", "-c", command
|
||||
], stdout=sys.stdout, stderr=sys.stderr, check=False)
|
||||
|
||||
sys.exit(result.returncode)
|
||||
|
||||
@app.function(gpu="L40S:2", image=image, timeout=3600)
|
||||
def run_ssim_tests():
|
||||
"""Run SSIM tests on 2x L40S GPUs"""
|
||||
import subprocess
|
||||
import sys
|
||||
import os
|
||||
|
||||
os.chdir("/FastVideo")
|
||||
|
||||
command = """
|
||||
source /opt/venv/bin/activate &&
|
||||
pytest ./fastvideo/v1/tests/ssim -vs
|
||||
"""
|
||||
|
||||
result = subprocess.run([
|
||||
"/bin/bash", "-c", command
|
||||
], stdout=sys.stdout, stderr=sys.stderr, check=False)
|
||||
|
||||
sys.exit(result.returncode)
|
||||
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-T2V-1.3B-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"))
|
||||
|
||||
|
||||
# training
|
||||
NUM_GPUS_PER_NODE_TRAINING = "4"
|
||||
TRAINING_ENTRY_FILE_PATH = "fastvideo/v1/training/wan_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_prompt_1_sample.json"),
|
||||
"--samples_per_file", "1",
|
||||
"--flush_frequency", "1",
|
||||
"--video_length_tolerance_range", "5",
|
||||
"--dataset", "t2v",
|
||||
]
|
||||
|
||||
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", "4",
|
||||
"--tp_size", "4",
|
||||
"--hsdp_replicate_dim", "1",
|
||||
"--hsdp_shard_dim", "4",
|
||||
"--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", "50",
|
||||
"--log_validation",
|
||||
"--checkpoints_total_limit", "3",
|
||||
"--allow_tf32",
|
||||
"--ema_start_step", "0",
|
||||
"--cfg", "0.0",
|
||||
"--output_dir", LOCAL_OUTPUT_DIR,
|
||||
"--tracker_project_name", "wan_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()
|
||||
@@ -0,0 +1,51 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
DATA_DIR="data/crush-smol_processed_main_t2v/latents/combined_parquet_dataset"
|
||||
VALIDATION_DIR="data/crush-smol_processed_main_t2v/latents/validation_parquet_dataset"
|
||||
NUM_GPUS=4
|
||||
# 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
|
||||
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
|
||||
fastvideo/v1/training/wan_training_pipeline.py\
|
||||
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
|
||||
--data_path "$DATA_DIR"\
|
||||
--validation_preprocessed_path "$VALIDATION_DIR"\
|
||||
--train_batch_size=1 \
|
||||
--num_latent_t 8 \
|
||||
--sp_size 4 \
|
||||
--tp_size 4 \
|
||||
--hsdp_replicate_dim 1 \
|
||||
--hsdp_shard_dim 4 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 1\
|
||||
--gradient_accumulation_steps=8 \
|
||||
--max_train_steps=5000 \
|
||||
--learning_rate=1e-5\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=6000 \
|
||||
--validation_steps 50\
|
||||
--validation_sampling_steps "50" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--output_dir="$DATA_DIR/outputs/wan_finetune"\
|
||||
--tracker_project_name wan_finetune \
|
||||
--num_height 480 \
|
||||
--num_width 832 \
|
||||
--num_frames 77 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--weight_decay 0.01 \
|
||||
--not_apply_cfg_solver \
|
||||
--dit_precision "fp32" \
|
||||
--max_grad_norm 1.0
|
||||
@@ -0,0 +1,24 @@
|
||||
# export WANDB_MODE="offline"
|
||||
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_main_t2v/latents"
|
||||
VALIDATION_PATH="examples/training/finetune/wan_t2v_1.3b/crush_smol/validation.json"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 77 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--model_type $MODEL_TYPE \
|
||||
--train_fps 16 \
|
||||
--validation_dataset_file $VALIDATION_PATH \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--preprocess_task "t2v"
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -78,7 +78,7 @@ I2V_MODEL_TO_PARAMS = {
|
||||
|
||||
TEST_PROMPTS = [
|
||||
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.",
|
||||
"A lone hiker stands atop a towering cliff, silhouetted against the vast horizon. The rugged landscape stretches endlessly beneath, its earthy tones blending into the soft blues of the sky. The scene captures the spirit of exploration and human resilience. High angle, dynamic framing, with soft natural lighting emphasizing the grandeur of nature."
|
||||
# "A lone hiker stands atop a towering cliff, silhouetted against the vast horizon. The rugged landscape stretches endlessly beneath, its earthy tones blending into the soft blues of the sky. The scene captures the spirit of exploration and human resilience. High angle, dynamic framing, with soft natural lighting emphasizing the grandeur of nature."
|
||||
]
|
||||
|
||||
I2V_TEST_PROMPTS = [
|
||||
@@ -328,5 +328,5 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
|
||||
if not success:
|
||||
logger.error("Failed to write SSIM results to file")
|
||||
|
||||
min_acceptable_ssim = 0.97
|
||||
min_acceptable_ssim = 0.95
|
||||
assert mean_ssim >= min_acceptable_ssim, f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim}"
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user