Compare commits

...
44 Commits
Author SHA1 Message Date
SolitaryThinker ea8c154176 ckpt 2025-06-28 14:56:21 -07:00
Kevin Lin 580d6dfe1f [CI] Add tests to Modal (#562) 2025-06-28 14:02:16 -05:00
Wenxuan Tan 344e43006a [CI] Fix SSIM and transformers CI (#564) 2025-06-28 00:26:20 -05:00
Wenxuan Tan c5155b256e [Feature] Load weights from distributed (#470) 2025-06-27 22:52:40 -05:00
William Lin e005c7f3ac [Docs] [Training] add readme for example training (#563) 2025-06-27 14:50:42 -05:00
Yongqi Chen ff5a79ef60 [Feature][Inference] Add VSA inference script (#561) 2025-06-27 02:19:23 -05:00
William Lin ab01dc4ba5 [Feature] [Training] Add i2v training (#559) 2025-06-27 01:56:50 -05:00
William Lin 285a950c1b [CI] fix vae and ssim tests (#557) 2025-06-26 23:53:01 -05:00
William Lin 46a0a85d85 [Training] Fixes SP for training; Improve Datasets and schema (#555) 2025-06-26 21:13:28 -05:00
Yongqi Chen 4aeabbc629 [Feature][Training] Add cfg rate for dataset loader (#556) 2025-06-26 18:22:37 -04:00
Wenxuan Tan 949bb5c835 [CI] Fix CI checks (#553) 2025-06-25 14:07:51 -05:00
Wenxuan Tan aab74c1271 [Kernel] Remove all syncs from STA & VSA kernels (#517) 2025-06-23 13:13:09 -07:00
Yongqi Chen f89d86944f [Feature][Training]Add diffusers format checkpoint saving for inference (#542) 2025-06-22 01:23:41 -04:00
William Lin 8741d204a5 [Training] Refactor and improve validation datasets (#539) 2025-06-21 17:58:35 -07:00
Wenxuan Tan cdc85f58a8 [chore] Bump torch to 2.7.1 to support Blackwell (#483) 2025-06-20 22:10:56 -07:00
William Lin 0262d2f089 [misc] [training] Reorganize training pipeline (#533) 2025-06-20 20:42:25 -07:00
William Lin 62c0343465 [bugfix] [VSA] Fix layernorm type for VSA Wan2.1 TransformerBlock (#534) 2025-06-20 00:24:51 -07:00
William Lin 1e1a023fb0 [bugfix] Fix stage validator for multi text encoder models (#535) 2025-06-19 22:49:16 -07:00
William Lin 1d2517ad8e [misc] Remove gradient checking code (#532) 2025-06-18 23:29:25 -07:00
William Lin d41186cb4a [Feat] Add Stage input and output verification (#523) 2025-06-18 23:29:11 -07:00
78e0c7eec9 Specify cu128 Pytorch installation (#530)
Co-authored-by: Edenzzzz <wtan45@wisc.edu>
Co-authored-by: Wenxuan Tan <wenxuan.tan@wisc.edu>
2025-06-18 20:02:50 -05:00
Wenxuan Tan 1c41a94b62 [Refactor] Move dict_to_3d_list under utils (#507) 2025-06-18 13:34:37 -07:00
Yongqi Chen 2e66aafe20 [Bugfix][Readme]Fix readme website bugs and add VSA finetune docs (#531) 2025-06-17 22:48:29 -07:00
Yongqi ChenandWill Lin 55074bda76 [CI] Add STA-inference/VSA-training test (#527)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-06-17 21:13:06 -07:00
William Lin de65bec2b7 [Ci] add sta and vsa install to docker image (#528) 2025-06-17 18:09:48 -07:00
Yongqi Chen 7664dd0de3 [Bugfix][Inference]Fix envs.attn_backend (#525) 2025-06-17 18:38:06 -05:00
William Linandkevin314 019a88ced4 [CI][bugfix] Use new 3.12 docker image (#526)
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
2025-06-17 15:37:08 -07:00
Kevin Lin 72de11abcc [CI] Add current PR test workflow to Buildkite/Modal (#512) 2025-06-17 13:29:22 -07:00
Kevin Lin d71a4ebffc [CI] Update Docker image to flash-attn 2.8.0 / CUDA 12.8 (#524) 2025-06-16 17:48:23 -07:00
William Lin 1089ab43bf [bugfix] [Training] use diffusers fp32layernorm for wan2.1 (#490) 2025-06-15 22:45:48 -07:00
William Lin 97d4b984c9 [misc] [ci] fix e2e preprocess+training data path (#521) 2025-06-14 22:37:51 -07:00
Wenxuan Tan 2a8953d74d [Refactor] Fix attn backend selection not correctly setting env variable (#516) 2025-06-15 00:04:54 -05:00
Yongqi Chen 8801b10da7 [Bugfix][Preprocess]fix mini dataset name (#520) 2025-06-14 22:03:22 -07:00
William Lin 6b413f2ec4 [CI] [Training] drop negative prompt in validation dataset and CI test for preprocess + training overfit (#519) 2025-06-14 18:50:17 -07:00
Yongqi Chen 28b72694aa [Feature][Preprocess]Add Readme doc for preprocess (#518) 2025-06-14 20:41:13 -04:00
Yongqi Chen 4afb0cfe4f [Feature][Training]vsa for t2v training ready (#513) 2025-06-14 01:08:00 -04:00
Zhang Peiyuan 3eec1281cf [misc] Fix preprocessing and dataloader extra padding (#514) 2025-06-13 15:15:33 -07:00
Wenxuan Tan 0660489e38 [CI] Restrict training CI to v1 (#508) 2025-06-12 15:26:05 -07:00
Zhang Peiyuan dd871a17bf fix logging (#509) 2025-06-12 15:24:12 -07:00
dc11529862 [Refactor][Configurations] clean config orgnization (#505)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-06-12 13:27:08 -07:00
Zhang Peiyuan ffabf85e31 [feat] Add parquet iterable dataset. (#506) 2025-06-12 04:30:56 -04:00
William Lin c0026ca5ba [CI] [Training] Initial e2e small training test (#504) 2025-06-11 13:53:36 -07:00
Zhang Peiyuan 0f2bbe71ac [misc] rename dp_size to hdsp_replicate_dim (#491) 2025-06-10 16:36:56 -07:00
Yongqi ChenandJerryZhou54 2a46902ecb [Feature][VSA]Update STA publish workflow (#498)
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
2025-06-10 19:33:34 -04:00
211 changed files with 8166 additions and 3174 deletions
+148
View File
@@ -0,0 +1,148 @@
env:
IMAGE_VERSION: "py3.12-latest"
steps:
- label: "pre-commit"
command: ".buildkite/scripts/pre_commit.sh"
agents:
queue: "default"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- wait
- label: "Trigger Tests"
plugins:
- monorepo-diff#v1.4.0:
diff: "git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD"
watch:
- path:
- "fastvideo/v1/models/encoders/**"
- "fastvideo/v1/models/loader/**"
- "fastvideo/v1/tests/encoders/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Encoder Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=encoder
agents:
queue: "default"
- path:
- "fastvideo/v1/models/vaes/**"
- "fastvideo/v1/models/loader/**"
- "fastvideo/v1/tests/vaes/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "VAE Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=vae
agents:
queue: "default"
- path:
- "fastvideo/v1/models/dits/**"
- "fastvideo/v1/models/loader/**"
- "fastvideo/v1/tests/transformers/**"
- "fastvideo/v1/layers/**"
- "fastvideo/v1/attention/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Transformer Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=transformer
agents:
queue: "default"
- path:
- "fastvideo/v1/**/*.py"
config:
command: "timeout 60m .buildkite/scripts/pr_test.sh"
label: "SSIM Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=ssim
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Training Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=training
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "csrc/attn/vsa/**"
- "csrc/attn/tk/**"
- "csrc/attn/setup_vsa.py"
- "csrc/attn/config_vsa.py"
- "csrc/attn/vsa.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Training Tests VSA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=training_vsa
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "csrc/attn/st_attn/**"
- "csrc/attn/setup_sta.py"
- "csrc/attn/config_sta.py"
- "csrc/attn/st_attn.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Inference Tests STA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=inference_sta
agents:
queue: "default"
- path:
- "csrc/attn/st_attn/**"
- "csrc/attn/setup_sta.py"
- "csrc/attn/config_sta.py"
- "csrc/attn/st_attn.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Precision Tests STA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=precision_sta
agents:
queue: "default"
- path:
- "csrc/attn/vsa/**"
- "csrc/attn/tk/**"
- "csrc/attn/setup_vsa.py"
- "csrc/attn/config_vsa.py"
- "csrc/attn/vsa.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VSA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=precision_vsa
agents:
queue: "default"
+117
View File
@@ -0,0 +1,117 @@
#!/bin/bash
set -uo pipefail
log() {
echo "[$(date '+%Y-%m-%d %H:%M:%S')] $1"
}
log "=== Starting Modal test execution ==="
# Change to the project directory
cd "$(dirname "$0")/../.."
PROJECT_ROOT=$(pwd)
log "Project root: $PROJECT_ROOT"
# Install Modal if not available
if ! python3 -m modal --version &> /dev/null; then
log "Modal not found, installing..."
python3 -m pip install modal
# Verify installation
if ! python3 -m modal --version &> /dev/null; then
log "Error: Failed to install modal. Please install it manually."
exit 1
fi
fi
log "modal version: $(python3 -m modal --version)"
# Set up Modal authentication using Buildkite secrets
log "Setting up Modal authentication from Buildkite secrets..."
MODAL_TOKEN_ID=$(buildkite-agent secret get modal_token_id)
MODAL_TOKEN_SECRET=$(buildkite-agent secret get modal_token_secret)
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
WANDB_API_KEY=$(buildkite-agent secret get wandb_api_key)
if [ -n "$MODAL_TOKEN_ID" ] && [ -n "$MODAL_TOKEN_SECRET" ]; then
log "Retrieved Modal credentials from Buildkite secrets"
python3 -m modal token set --token-id "$MODAL_TOKEN_ID" --token-secret "$MODAL_TOKEN_SECRET" --profile buildkite-ci --activate --verify
if [ $? -eq 0 ]; then
log "Modal authentication successful"
else
log "Error: Failed to set Modal credentials"
exit 1
fi
else
log "Error: Could not retrieve Modal credentials from Buildkite secrets."
log "Please ensure 'modal_token_id' and 'modal_token_secret' secrets are set in Buildkite."
exit 1
fi
MODAL_TEST_FILE="fastvideo/v1/tests/modal/pr_test.py"
if [ -z "${TEST_TYPE:-}" ]; then
log "Error: TEST_TYPE environment variable is not set"
exit 1
fi
log "Test type: $TEST_TYPE"
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT IMAGE_VERSION=$IMAGE_VERSION"
case "$TEST_TYPE" in
"encoder")
log "Running encoder tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
;;
"vae")
log "Running VAE tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
;;
"transformer")
log "Running transformer tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
;;
"ssim")
log "Running SSIM tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
;;
"training")
log "Running training tests..."
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests"
;;
"training_vsa")
log "Running training VSA tests..."
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_training_tests_VSA"
;;
"inference_sta")
log "Running inference STA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_STA"
;;
"precision_sta")
log "Running precision STA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_STA"
;;
"precision_vsa")
log "Running precision VSA tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_VSA"
;;
*)
log "Error: Unknown test type: $TEST_TYPE"
exit 1
;;
esac
log "Executing: $MODAL_COMMAND"
eval "$MODAL_COMMAND"
TEST_EXIT_CODE=$?
if [ $TEST_EXIT_CODE -eq 0 ]; then
log "Modal test completed successfully"
else
log "Error: Modal test failed with exit code: $TEST_EXIT_CODE"
fi
log "=== Test execution completed with exit code: $TEST_EXIT_CODE ==="
exit $TEST_EXIT_CODE
+40
View File
@@ -0,0 +1,40 @@
#!/bin/bash
set -uo pipefail
log() {
echo "[$(date '+%Y-%m-%d %H:%M:%S')] $1"
}
log "=== Starting pre-commit checks ==="
cd "$(dirname "$0")/../.."
PROJECT_ROOT=$(pwd)
log "Project root: $PROJECT_ROOT"
if ! python3 -m pre_commit --version &> /dev/null; then
log "pre-commit not found, installing..."
python3 -m pip install --user pre-commit==4.0.1
if ! python3 -m pre_commit --version &> /dev/null; then
log "Error: Failed to install pre-commit."
exit 1
fi
fi
log "Pre-commit version: $(python3 -m pre_commit --version)"
log "Installing/updating pre-commit hooks..."
python3 -m pre_commit install --install-hooks
log "Running pre-commit checks on all files..."
python3 -m pre_commit run --all-files
PRE_COMMIT_EXIT_CODE=$?
if [ $PRE_COMMIT_EXIT_CODE -eq 0 ]; then
log "Pre-commit checks completed successfully"
else
log "Error: Pre-commit checks failed with exit code: $PRE_COMMIT_EXIT_CODE"
fi
log "=== Pre-commit checks completed with exit code: $PRE_COMMIT_EXIT_CODE ==="
exit $PRE_COMMIT_EXIT_CODE
+1 -2
View File
@@ -160,8 +160,7 @@ def execute_command(pod_id):
setup_steps = [
"tar -xzf /tmp/repo.tar.gz --no-same-owner -C /workspace/",
f"cd /workspace/{repo_name}",
"source /opt/conda/etc/profile.d/conda.sh",
"conda activate fastvideo-dev",
"source $HOME/.local/bin/env && source /opt/venv/bin/activate",
args.test_command
]
+209 -19
View File
@@ -12,13 +12,11 @@ on:
paths:
- "fastvideo/**/*.py"
- ".github/workflows/pr-test.yml"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
- "csrc/**"
workflow_dispatch:
inputs:
custom_image:
description: "Custom image from this repository (default: fastvideo-dev:latest)"
required: false
default: "fastvideo-dev:latest"
type: string
run_encoder_test:
description: "Run encoder-test"
required: false
@@ -39,10 +37,41 @@ on:
required: false
default: false
type: boolean
run_training_test:
description: "Run training-test"
required: false
default: false
type: boolean
run_training_test_VSA:
description: "Run training-test-VSA"
required: false
default: false
type: boolean
run_inference_test_STA:
description: "Run inference-test-STA"
required: false
default: false
type: boolean
run_precision_test_STA:
description: "Run precision-test-STA"
required: false
default: false
type: boolean
run_precision_test_VSA:
description: "Run precision-test-VSA"
required: false
default: false
type: boolean
run_nightly_test:
description: "Run nightly-test"
required: false
default: false
type: boolean
env:
PYTHONUNBUFFERED: "1"
concurrency:
group: pr-test-${{ github.ref }}
cancel-in-progress: true
@@ -59,26 +88,72 @@ jobs:
encoder-test: ${{ steps.filter.outputs.encoder-test }}
vae-test: ${{ steps.filter.outputs.vae-test }}
transformer-test: ${{ steps.filter.outputs.transformer-test }}
training-test: ${{ steps.filter.outputs.training-test }}
training-test-VSA: ${{ steps.filter.outputs.training-test-VSA }}
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
steps:
- uses: actions/checkout@v4
- uses: dorny/paths-filter@v3
id: filter
with:
filters: |
# Define reusable path patterns
common-paths: &common-paths
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
sta-kernel-paths: &sta-kernel-paths
- 'csrc/attn/st_attn/**'
- 'csrc/attn/setup_sta.py'
- 'csrc/attn/config_sta.py'
- 'csrc/attn/st_attn.cpp'
vsa-kernel-paths: &vsa-kernel-paths
- 'csrc/attn/vsa/**'
- 'csrc/attn/tk/**'
- 'csrc/attn/setup_vsa.py'
- 'csrc/attn/config_vsa.py'
- 'csrc/attn/vsa.cpp'
vsa-paths: &vsa-paths
- 'fastvideo/v1/**'
- *common-paths
- *vsa-kernel-paths
# Actual tests
encoder-test:
- 'fastvideo/v1/models/encoders/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/models/loader/**'
- 'fastvideo/v1/tests/encoders/**'
- *common-paths
vae-test:
- 'fastvideo/v1/models/vaes/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/models/loader/**'
- 'fastvideo/v1/tests/vaes/**'
- *common-paths
transformer-test:
- 'fastvideo/v1/models/dits/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/models/loader/**'
- 'fastvideo/v1/tests/transformers/**'
- 'fastvideo/v1/layers/**'
- 'fastvideo/v1/attention/**'
- *common-paths
training-test:
- 'fastvideo/v1/**'
- *common-paths
training-test-VSA:
- 'fastvideo/v1/**'
- *common-paths
- *vsa-kernel-paths
inference-test-STA:
- 'fastvideo/v1/**'
- *common-paths
- *sta-kernel-paths
precision-test-STA:
- *common-paths
- *sta-kernel-paths
precision-test-VSA:
- *common-paths
- *vsa-kernel-paths
encoder-test:
needs: change-filter
@@ -91,8 +166,8 @@ jobs:
gpu_type: "NVIDIA A40"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -109,8 +184,8 @@ jobs:
gpu_type: "NVIDIA A40"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -127,8 +202,8 @@ jobs:
gpu_type: "NVIDIA L40S"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
@@ -137,8 +212,7 @@ jobs:
ssim-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
github.event_name != 'workflow_dispatch' || (github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
strategy:
fail-fast: false
matrix:
@@ -155,14 +229,130 @@ jobs:
volume_size: 200
disk_size: 200
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:${{ matrix.python-version.tag }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
timeout_minutes: 60
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
training-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "training-test"
gpu_type: "NVIDIA A40"
gpu_count: 4
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/Vanilla -srP"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
training-test-VSA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.training-test-VSA == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test_VSA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "training-test-VSA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/training/VSA -srP"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
inference-test-STA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.inference-test-STA == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_inference_test_STA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "inference-test-STA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/inference/STA -srP"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
precision-test-STA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-STA == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_STA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "precision-test-STA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_sta.py"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
precision-test-VSA:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.precision-test-VSA == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_precision_test_VSA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "precision-test-VSA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "uv pip install -e .[test] && python csrc/attn/tests/test_block_sparse.py"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
nightly-test:
if: >-
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "nightly-test"
gpu_type: "NVIDIA A40"
gpu_count: 4
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/v1/tests/nightly/test_e2e_overfit_single_sample.py -vs"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
runpod-cleanup:
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
# Add other jobs to this list as you create them
needs: [encoder-test, vae-test, transformer-test, ssim-test, training-test, training-test-VSA, inference-test-STA, precision-test-STA, precision-test-VSA]
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
runs-on: ubuntu-latest
steps:
@@ -179,7 +369,7 @@ jobs:
- name: Cleanup all RunPod instances
env:
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12"]'
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12", "training-test", "training-test-VSA", "inference-test-STA", "precision-test-STA", "precision-test-VSA"]'
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
run: python .github/scripts/runpod_cleanup.py
+4 -1
View File
@@ -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
+17 -13
View File
@@ -5,7 +5,7 @@ on:
branches:
- main
paths:
- "csrc/sliding_tile_attention/setup.py"
- "csrc/attn/setup_sta.py"
workflow_dispatch:
jobs:
@@ -23,13 +23,13 @@ jobs:
- name: Check if version changed
id: check-version
run: |
cd csrc/sliding_tile_attention
cd csrc/attn
# Get current commit's version
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_sta.py)
echo "New version: $NEW_VERSION"
# Get previous version from git history
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
OLD_VERSION=$(git show HEAD~1:./setup_sta.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
@@ -136,19 +136,21 @@ jobs:
- name: Build wheel
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
# However this still fails so I'm using a newer version of setuptools
pip install setuptools
pip install ninja packaging wheel
cd csrc/sliding_tile_attention # Move into the correct folder
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
python setup.py bdist_wheel --dist-dir=dist
cd csrc/attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup_sta.py bdist_wheel --dist-dir=dist
- name: Rename wheel file
run: |
cd csrc/sliding_tile_attention
cd csrc/attn
CUDA_SHORT_VERSION=$(echo ${{ matrix.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-version }} | cut -d. -f1,2)
@@ -163,7 +165,7 @@ jobs:
uses: actions/upload-artifact@v4
with:
name: ${{ env.wheel_name }}
path: csrc/sliding_tile_attention/dist/*.whl
path: csrc/attn/dist/*.whl
retention-days: 90
publish_package:
@@ -229,17 +231,19 @@ jobs:
- name: Build source distribution
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
# However this still fails so I'm using a newer version of setuptools
pip install setuptools
pip install ninja packaging wheel
cd csrc/sliding_tile_attention # Move into the correct folder
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
python setup.py sdist --dist-dir=dist
cd csrc/attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup_sta.py sdist --dist-dir=dist
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: csrc/sliding_tile_attention/dist/
packages-dir: csrc/attn/dist/
+1 -1
View File
@@ -28,4 +28,4 @@ jobs:
- name: Run Pytest
run: |
pytest --ignore csrc/sliding_tile_attention/test
pytest --ignore csrc/attn/test
-1
View File
@@ -27,7 +27,6 @@ env
**/build/
**.pyc
**.txt
csrc/attn/tk/
# Distribution / packaging
build/
+2 -2
View File
@@ -1,3 +1,3 @@
[submodule "csrc/sliding_tile_attention/tk"]
path = csrc/sliding_tile_attention/tk
[submodule "csrc/attn/tk"]
path = csrc/attn/tk
url = https://github.com/HazyResearch/ThunderKittens.git
+2 -2
View File
@@ -91,7 +91,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
## Distillation and Finetuning
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/training/distillation.html)
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetuning.html)
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html)
## 📑 Development Plan
@@ -111,7 +111,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
## 🤝 Contributing
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/developer_guide/overview.html)
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview.html)
## Acknowledgement
We learned and reused code from the following projects:
+21 -6
View File
@@ -4,9 +4,8 @@
## Installation
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only support H100/H200, because ThunderKittens uses TMA but doesn't support Blackwell yet.
First, install C++20 for ThunderKittens:
```bash
sudo apt update
sudo apt install gcc-11 g++-11
@@ -16,17 +15,27 @@ sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave
sudo apt update
sudo apt install clang-11
```
Install STA:
## Environment Setup
First, set up your CUDA environment:
```bash
export CUDA_HOME=/usr/local/cuda-12.4
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
git submodule update --init --recursive
python setup.py install
```
## Install Sliding Tile Attention (STA)
```bash
python setup_sta.py install
```
## Install Video Sparse Attention (VSA)
```bash
python setup_vsa.py install
```
## Usage
```python
from st_attn import sliding_tile_attention
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
@@ -44,8 +53,14 @@ out = sliding_tile_attention(q, k, v, window_size, 0, False)
## Test
```bash
python test/test_sta.py
python tests/test_sta.py # test STA
python tests/test_block_sparse.py # test VSA
```
## Benchmark
```bash
python benchmarks/bench_sta.py
```
## How Does STA Work?
We give a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
@@ -5,6 +5,7 @@ import matplotlib.pyplot as plt
import numpy as np
import torch
from st_attn import sliding_tile_attention
from triton.testing import do_bench
def flops(batch, seqlen, nheads, headdim, causal, mode="fwd"):
@@ -13,16 +14,16 @@ def flops(batch, seqlen, nheads, headdim, causal, mode="fwd"):
return f if mode == "fwd" else (2.5 * f if mode == "bwd" else 3.5 * f)
def efficiency(flop, time):
flop = flop / 1e12
time = time / 1e6
return flop / time
def compute_TFLOPS(flops, ms):
flops = flops / 1e12
ms = ms / 1e3
return flops / ms
def benchmark_attention(configurations):
results = {'fwd': defaultdict(list), 'bwd': defaultdict(list)}
for B, H, N, D, causal in configurations:
for B, H, N, D, causal, dit_seq_shape, window_size in configurations:
print("=" * 60)
print(f"Timing forward and backward pass for B={B}, H={H}, N={N}, D={D}, causal={causal}")
@@ -30,38 +31,31 @@ def benchmark_attention(configurations):
k = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
v = torch.randn(B, H, N, D, dtype=torch.bfloat16, device='cuda', requires_grad=False).contiguous()
grad_output = torch.randn_like(q, requires_grad=False).contiguous()
# grad_output = torch.randn_like(q, requires_grad=False).contiguous()
# qg = torch.zeros_like(q, requires_grad=False, dtype=torch.float).contiguous()
# kg = torch.zeros_like(k, requires_grad=False, dtype=torch.float).contiguous()
# vg = torch.zeros_like(v, requires_grad=False, dtype=torch.float).contiguous()
qg = torch.zeros_like(q, requires_grad=False, dtype=torch.float).contiguous()
kg = torch.zeros_like(k, requires_grad=False, dtype=torch.float).contiguous()
vg = torch.zeros_like(v, requires_grad=False, dtype=torch.float).contiguous()
# Prepare for timing forward pass
start_events_fwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
end_events_fwd = [torch.cuda.Event(enable_timing=True) for _ in range(10)]
# # Warmup for forward pass
# for _ in range(10):
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
torch.cuda.empty_cache()
torch.cuda.synchronize()
# # Time the forward pass
# for i in range(10):
# start_events_fwd[i].record()
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
# end_events_fwd[i].record()
ms = do_bench(lambda: sliding_tile_attention(q, k, v, [window_size] * 24, 0, False, dit_seq_shape))
# Warmup for forward pass
for _ in range(10):
o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, '18x48x80')
# times_fwd = [s.elapsed_time(e) for s, e in zip(start_events_fwd, end_events_fwd)]
# time_us_fwd = np.mean(times_fwd) * 1000
# Time the forward pass
for i in range(10):
start_events_fwd[i].record()
o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, '18x48x80')
end_events_fwd[i].record()
torch.cuda.synchronize()
times_fwd = [s.elapsed_time(e) for s, e in zip(start_events_fwd, end_events_fwd)]
time_us_fwd = np.mean(times_fwd) * 1000
tflops_fwd = efficiency(flops(B, N, H, D, causal, 'fwd'), time_us_fwd)
tflops_fwd = compute_TFLOPS(flops(B, N, H, D, causal, 'fwd'), ms)
results['fwd'][(D, causal)].append((N, tflops_fwd))
print(f"Average time for forward pass in us: {time_us_fwd:.2f}")
print(f"Average efficiency for forward pass in TFLOPS: {tflops_fwd}")
print(f"Average time for forward pass (ms): {ms:.2f}")
print(f"Average TFLOPS: {tflops_fwd}")
print("-" * 60)
# torch.cuda.empty_cache()
@@ -85,15 +79,14 @@ def benchmark_attention(configurations):
# times_bwd = [s.elapsed_time(e) for s, e in zip(start_events_bwd, end_events_bwd)]
# time_us_bwd = np.mean(times_bwd) * 1000
# tflops_bwd = efficiency(flops(B, N, H, D, causal, 'bwd'), time_us_bwd)
# tflops_bwd = compute_TFLOPS(flops(B, N, H, D, causal, 'bwd'), ms)
# results['bwd'][(D, causal)].append((N, tflops_bwd))
# print(f"Average time for backward pass in us: {time_us_bwd:.2f}")
# print(f"Average efficiency for backward pass in TFLOPS: {tflops_bwd}")
print("=" * 60)
# print(f"Average time for backward pass(ms): {ms:.2f}")
# print(f"Average TFLOPS: {tflops_bwd}")
# print("=" * 60)
torch.cuda.empty_cache()
torch.cuda.synchronize()
return results
@@ -124,7 +117,10 @@ def plot_results(results):
# Example list of configurations to test
configurations = [
(2, 24, 69120, 128, False),
(2, 24, 69120, 128, False, '18x48x80', [3, 6, 10]),
(2, 24, 69120, 128, True, '18x48x80', [3, 6, 10]),
(2, 24, 82944, 128, False, '36x48x48', [3, 3, 6]), # Stepvideo
(2, 24, 82944, 128, True, '36x48x48', [3, 3, 6]),
# (16, 16, 768*16, 128, False),
# (16, 16, 768*2, 128, False),
# (16, 16, 768*4, 128, False),
+31 -22
View File
@@ -4,9 +4,17 @@
#include <cooperative_groups.h>
#include <iostream>
#include <stdio.h>
#include <c10/cuda/CUDAGuard.h>
// #define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
__device__ __forceinline__ int clamp_int(int value, int min, int max) {
return (value < min) ? min : ((value > max) ? max : value);
}
// #define ABS(x) ((x) < 0 ? -(x) : (x))
__device__ __forceinline__ int abs_int(int value) {
return (value < 0) ? -value : value;
}
#define CLAMP(value, min, max) ((value) < (min) ? (min) : ((value) > (max) ? (max) : (value)))
#define ABS(x) ((x) < 0 ? -(x) : (x))
constexpr int CONSUMER_WARPGROUPS = (3);
constexpr int PRODUCER_WARPGROUPS = (1);
@@ -117,16 +125,16 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
int qt = seq_idx / 6 / (CH * CW);
int qh = (seq_idx / 6) % (CH * CW) / CW;
int qw = (seq_idx / 6) % CW;
qt = CLAMP(qt, DT, CT-DT-1);
qh = CLAMP(qh, DH, CH-DH-1);
qw = CLAMP(qw, DW, CW-DW-1);
qt = clamp_int(qt, DT, CT-DT-1);
qh = clamp_int(qh, DH, CH-DH-1);
qw = clamp_int(qw, DW, CW-DW-1);
int count = 0;
int j = 0;
while (count < K::stages - 1) {
int kt = j / 3 / (CH * CW);
int kh = (j / 3) % (CH * CW) / CW;
int kw = (j / 3) % CW;
bool mask = (ABS(qt - kt) <= DT) && (ABS(qh - kh) <= DH) && (ABS(qw - kw) <= DW);
bool mask = (abs_int(qt - kt) <= DT) && (abs_int(qh - kh) <= DH) && (abs_int(qw - kw) <= DW);
if (mask){
coord<k_tile> kv_tile_idx = {blockIdx.z, kv_head_idx, j, 0};
tma::expect_bytes(k_smem_arrived[count], sizeof(k_tile));
@@ -167,15 +175,15 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
int qt = seq_idx / 6 / (CH * CW);
int qh = (seq_idx / 6) % (CH * CW) / CW;
int qw = (seq_idx / 6) % CW;
qt = CLAMP(qt, DT, CT-DT-1);
qh = CLAMP(qh, DH, CH-DH-1);
qw = CLAMP(qw, DW, CW-DW-1);
int k_t_min = CLAMP(qt-DT, 0, CT-1);
int k_t_max = CLAMP(qt+DT, 0, CT-1);
int k_h_min = CLAMP(qh-DH, 0, CH-1);
int k_h_max = CLAMP(qh+DH, 0, CH-1);
int k_w_min = CLAMP(qw-DW, 0, CW-1);
int k_w_max = CLAMP(qw+DW, 0, CW-1);
qt = clamp_int(qt, DT, CT-DT-1);
qh = clamp_int(qh, DH, CH-DH-1);
qw = clamp_int(qw, DW, CW-DW-1);
int k_t_min = clamp_int(qt-DT, 0, CT-1);
int k_t_max = clamp_int(qt+DT, 0, CT-1);
int k_h_min = clamp_int(qh-DH, 0, CH-1);
int k_h_max = clamp_int(qh+DH, 0, CH-1);
int k_w_min = clamp_int(qw-DW, 0, CW-1);
int k_w_max = clamp_int(qw+DW, 0, CW-1);
int count = 0;
for (int kt = k_t_min; kt <= k_t_max; kt++) {
for (int kh = k_h_min; kh <= k_h_max; kh++) {
@@ -234,7 +242,7 @@ void fwd_attend_ker(const __grid_constant__ fwd_globals<D> g) {
// the last three kv blocks are for text, we process them separately
kv_iters = img_kv_blocks - 1;
} else {
kv_iters = CLAMP(DT*2+1, 1, CT) * CLAMP(DH*2+1, 1, CH) * CLAMP(DW*2+1, 1, CW) * 3 - 1 ;
kv_iters = clamp_int(DT*2+1, 1, CT) * clamp_int(DH*2+1, 1, CH) * clamp_int(DW*2+1, 1, CW) * 3 - 1 ;
}
kittens::wait(qsmem_semaphore, 0);
@@ -415,8 +423,9 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
float* l_ptr = reinterpret_cast<float*>(l_vec.data_ptr<float>());
float* d_l = reinterpret_cast<float*>(l_ptr);
cudaDeviceSynchronize();
auto stream = at::cuda::getCurrentCUDAStream().stream();
//cudadevicesynchronize();
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
if (head_dim == 128) {
@@ -442,8 +451,8 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
globals g{qg_arg, kg_arg, vg_arg, lg_arg, og_arg, static_cast<int>(seq_len), static_cast<int>(text_length), static_cast<int>(hr)};
auto mem_size = kittens::MAX_SHARED_MEMORY;
auto threads = NUM_WORKERS * kittens::WARP_THREADS;
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
int threads = NUM_WORKERS * kittens::WARP_THREADS;
if (has_text) {
// TORCH_CHECK(seq_len % (CONSUMER_WARPGROUPS*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 192");
dim3 grid_image(seq_len/(CONSUMER_WARPGROUPS*kittens::TILE_ROW_DIM<bf16>*4)-2, qo_heads, batch);
@@ -823,10 +832,10 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
}
CHECK_CUDA_ERROR(cudaGetLastError());
cudaStreamSynchronize(stream);
// cudaStreamSynchronize(stream);
}
return o;
cudaDeviceSynchronize();
//cudadevicesynchronize();
}
@@ -2,11 +2,12 @@ import torch
import argparse
from flash_attn.utils.benchmark import benchmark_forward
from flash_attn import flash_attn_func
from st_attn import block_sparse_attention_fwd, block_sparse_attention_backward, BlockSparseAttentionFunction
from st_attn import BLOCK_M, BLOCK_N
from vsa import block_sparse_attention_fwd, block_sparse_attention_backward, BlockSparseAttentionFunction
from vsa import BLOCK_M, BLOCK_N
import numpy as np
import random
import gc
def set_seed(seed: int = 42):
# Python random module
@@ -20,15 +21,6 @@ def set_seed(seed: int = 42):
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed) # if using multi-GPU
def parse_arguments():
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
parser.add_argument('--batch_size', type=int, default=4, help='Batch size')
parser.add_argument('--num_heads', type=int, default=6, help='Number of heads')
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
parser.add_argument('--topk', type=int, default=64, help='Number of kv blocks each q block attends to')
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[29120], help='Sequence lengths to benchmark')
parser.add_argument('--num_iterations', type=int, default=100, help='Number of test iterations to run')
return parser.parse_args()
@torch.no_grad
def precision_metric(quant_o, fa2_o):
@@ -135,9 +127,7 @@ def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device=
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
def main():
args = parse_arguments()
def main(args):
set_seed(42)
# Extract parameters
@@ -191,23 +181,36 @@ def main():
block_mask_expanded = block_sparse_mask.unsqueeze(-1).unsqueeze(-2) # [b, h, num_q_blocks, num_kv_blocks, 1, 1]
block_mask_expanded = block_mask_expanded.expand(-1, -1, -1, -1, BLOCK_M, BLOCK_N) # [b, h, num_q_blocks, num_kv_blocks, BLOCK_M, BLOCK_N]
full_mask = block_mask_expanded.permute(0, 1, 2, 4, 3, 5).reshape(batch, head, seq_len, seq_len)
q_sdpa = q.clone()
k_sdpa = k.clone()
v_sdpa = v.clone()
q.requires_grad = True
k.requires_grad = True
v.requires_grad = True
q_sdpa.requires_grad = True
k_sdpa.requires_grad = True
v_sdpa.requires_grad = True
# testing forward
o = BlockSparseAttentionFunction.apply(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
del q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask, block_mask_expanded
grad_o = torch.randn_like(o)
o.backward(grad_o)
# clear memory
q_sdpa = q.detach().clone()
k_sdpa = k.detach().clone()
v_sdpa = v.detach().clone()
q_sdpa.requires_grad = True
k_sdpa.requires_grad = True
v_sdpa.requires_grad = True
q.data = torch.empty(0, device=q.device)
k.data = torch.empty(0, device=k.device)
v.data = torch.empty(0, device=v.device)
torch.cuda.empty_cache()
o_sdpa = torch.nn.functional.scaled_dot_product_attention(q_sdpa, k_sdpa, v_sdpa, attn_mask=full_mask)
sim, l1, rmse = precision_metric(o, o_sdpa)
assert sim > 0.9999, f"SSIM too low: {sim}"
assert l1 < 8e-5, f"l1 too large: {l1}"
assert rmse < 2e-5, f"RMSE too large: {rmse}"
forward_metrics['sim'].append(sim)
forward_metrics['l1'].append(l1)
forward_metrics['rmse'].append(rmse)
@@ -215,52 +218,72 @@ def main():
print(f"block_sparse_attention_fwd vs torch.nn.functional.scaled_dot_product_attention:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
# test backward
grad_o = torch.randn_like(o)
o.backward(grad_o)
o_sdpa.backward(grad_o)
sim, l1, rmse = precision_metric(q.grad, q_sdpa.grad)
# Error bounds collected on H100
assert sim > 0.9999, f"SSIM too low: {sim}"
assert l1 < 4e-3, f"l1 too large: {l1}"
assert rmse < 3e-4, f"RMSE too large: {rmse}"
grad_q_metrics['sim'].append(sim)
grad_q_metrics['l1'].append(l1)
grad_q_metrics['rmse'].append(rmse)
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_q:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
sim, l1, rmse = precision_metric(k.grad, k_sdpa.grad)
assert sim > 0.9999, f"SSIM too low: {sim}"
assert l1 < 4e-3, f"l1 too large: {l1}"
assert rmse < 2e-4, f"RMSE too large: {rmse}"
grad_k_metrics['sim'].append(sim)
grad_k_metrics['l1'].append(l1)
grad_k_metrics['rmse'].append(rmse)
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_k:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
sim, l1, rmse = precision_metric(v.grad, v_sdpa.grad)
assert sim > 0.9999, f"SSIM too low: {sim}"
assert l1 < 1e-4, f"l1 too large: {l1}"
assert rmse < 2e-5, f"RMSE too large: {rmse}"
grad_v_metrics['sim'].append(sim)
grad_v_metrics['l1'].append(l1)
grad_v_metrics['rmse'].append(rmse)
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_v:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
del o, o_sdpa, grad_o, q_sdpa, k_sdpa, v_sdpa
gc.collect()
torch.cuda.empty_cache()
# Print summary statistics if multiple iterations were run
if num_iterations > 1:
print("\n" + "="*50)
print(f"Summary Statistics (over {num_iterations} iterations):")
print("\nForward metrics:")
print(f"Similarity: mean={np.mean(forward_metrics['sim']):.6f}, std={np.std(forward_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(forward_metrics['l1']):.6f}, std={np.std(forward_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(forward_metrics['rmse']):.6f}, std={np.std(forward_metrics['rmse']):.6f}")
print(f"Similarity: mean={np.mean(forward_metrics['sim']):.6f}, std={np.std(forward_metrics['sim']):.6f}, min={np.min(forward_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(forward_metrics['l1']):.6f}, std={np.std(forward_metrics['l1']):.6f}, max={np.max(forward_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(forward_metrics['rmse']):.6f}, std={np.std(forward_metrics['rmse']):.6f}, max={np.max(forward_metrics['rmse']):.6f}")
print("\nGradient Q metrics:")
print(f"Similarity: mean={np.mean(grad_q_metrics['sim']):.6f}, std={np.std(grad_q_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_q_metrics['l1']):.6f}, std={np.std(grad_q_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_q_metrics['rmse']):.6f}, std={np.std(grad_q_metrics['rmse']):.6f}")
print(f"Similarity: mean={np.mean(grad_q_metrics['sim']):.6f}, std={np.std(grad_q_metrics['sim']):.6f}, min={np.min(grad_q_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_q_metrics['l1']):.6f}, std={np.std(grad_q_metrics['l1']):.6f}, max={np.max(grad_q_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_q_metrics['rmse']):.6f}, std={np.std(grad_q_metrics['rmse']):.6f}, max={np.max(grad_q_metrics['rmse']):.6f}")
print("\nGradient K metrics:")
print(f"Similarity: mean={np.mean(grad_k_metrics['sim']):.6f}, std={np.std(grad_k_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_k_metrics['l1']):.6f}, std={np.std(grad_k_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_k_metrics['rmse']):.6f}, std={np.std(grad_k_metrics['rmse']):.6f}")
print(f"Similarity: mean={np.mean(grad_k_metrics['sim']):.6f}, std={np.std(grad_k_metrics['sim']):.6f}, min={np.min(grad_k_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_k_metrics['l1']):.6f}, std={np.std(grad_k_metrics['l1']):.6f}, max={np.max(grad_k_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_k_metrics['rmse']):.6f}, std={np.std(grad_k_metrics['rmse']):.6f}, max={np.max(grad_k_metrics['rmse']):.6f}")
print("\nGradient V metrics:")
print(f"Similarity: mean={np.mean(grad_v_metrics['sim']):.6f}, std={np.std(grad_v_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_v_metrics['l1']):.6f}, std={np.std(grad_v_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_v_metrics['rmse']):.6f}, std={np.std(grad_v_metrics['rmse']):.6f}")
print(f"Similarity: mean={np.mean(grad_v_metrics['sim']):.6f}, std={np.std(grad_v_metrics['sim']):.6f}, min={np.min(grad_v_metrics['sim']):.6f}")
print(f"L1 error: mean={np.mean(grad_v_metrics['l1']):.6f}, std={np.std(grad_v_metrics['l1']):.6f}, max={np.max(grad_v_metrics['l1']):.6f}")
print(f"RMSE: mean={np.mean(grad_v_metrics['rmse']):.6f}, std={np.std(grad_v_metrics['rmse']):.6f}, max={np.max(grad_v_metrics['rmse']):.6f}")
if __name__ == "__main__":
main()
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
parser.add_argument('--batch_size', type=int, default=4, help='Batch size')
parser.add_argument('--num_heads', type=int, default=6, help='Number of heads')
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
parser.add_argument('--topk', type=int, default=64, help='Number of kv blocks each q block attends to')
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[29120], help='Sequence lengths to benchmark')
parser.add_argument('--num_iterations', type=int, default=50, help='Number of test iterations to run')
args = parser.parse_args()
main(args)
@@ -9,14 +9,14 @@ flex_attention = torch.compile(flex_attention, dynamic=False)
def flex_test(Q, K, V, kernel_size):
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (36, 48, 48), 39, 'cuda', 0)
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (18, 48, 80), 0, 'cuda', 0)
output = flex_attention(Q, K, V, block_mask=mask)
return output
def h100_fwd_kernel_test(Q, K, V, kernel_size):
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 39, False)
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
return o
@@ -37,7 +37,7 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
'max_diff': 0
},
}
kernel_size_ls = [(6, 1, 6), (6, 6, 1)]
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
from tqdm import tqdm
for kernel_size in tqdm(kernel_size_ls):
for _ in range(num_iterations):
@@ -74,12 +74,14 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
# Example usage
b, h, d = 2, 24, 128
n = 82944 # Sequence length
n = 69120 # Sequence length
causal = False
mean = 1e-1
std = 10
# Run correctness check directly
results = check_correctness(b, h, n, d, causal, mean, std, error_mode='output')
assert results['TK vs FLEX']['avg_diff'] < 3e-6, f"Average difference: {results['TK vs FLEX']['avg_diff']} is too large"
assert results['TK vs FLEX']['max_diff'] < 4e-2, f"Maximum difference: {results['TK vs FLEX']['max_diff']} is too large"
print(f"Average difference: {results['TK vs FLEX']['avg_diff']}")
print(f"Maximum difference: {results['TK vs FLEX']['max_diff']}")
Submodule
+1
Submodule csrc/attn/tk added at 1719fb7264
+284 -24
View File
@@ -2,7 +2,7 @@ import math
import torch
from torch.utils.checkpoint import detach_variable
from typing import Tuple
try:
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
except ImportError:
@@ -12,33 +12,116 @@ except ImportError:
BLOCK_M = 64
BLOCK_N = 64
def block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num):
def video_sparse_attn(q, k, v, topk, block_size, compress_attn_weight=None):
"""
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks].
[*, *, i, j] = 1 means the i-th q block should attend to the j-th kv block.
q: [batch_size, num_heads, seq_len, head_dim]
k: [batch_size, num_heads, seq_len, head_dim]
v: [batch_size, num_heads, seq_len, head_dim]
topk: int
block_size: int or tuple of 3 ints
video_shape: tuple of (T, H, W)
compress_attn_weight: [batch_size, num_heads, seq_len, head_dim]
select_attn_weight: [batch_size, num_heads, seq_len, head_dim]
V1 of sparse attention. Include compress attn and sparse attn branch, use average pooling to compress.
Assume q, k, v is flattened in this way: [batch_size, num_heads, T//block_size[0], H//block_size[1], W//block_size[2], block_size[0], block_size[1], block_size[2]]
"""
# assert all elements in q2k_block_sparse_num can be devisible by 2
o, lse = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
return o, lse
def block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num):
grad_q, grad_k, grad_v = block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
return grad_q, grad_k, grad_v
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
class BlockSparseAttentionFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
o, lse = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
ctx.save_for_backward(q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num)
return o
block_elements = block_size[0] * block_size[1] * block_size[2]
assert block_elements % 64 == 0 and block_elements >= 64
assert q.shape[2] % block_elements == 0
batch_size, num_heads, seq_len, head_dim = q.shape
# compress attn
q_compress = q.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
k_compress = k.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
v_compress = v.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
@staticmethod
def backward(ctx, grad_output):
q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num = ctx.saved_tensors
grad_q, grad_k, grad_v = block_sparse_attention_backward(
q, k, v, o, lse, grad_output, k2q_block_sparse_index, k2q_block_sparse_num
)
return grad_q, grad_k, grad_v, None, None, None, None
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
v_compress)
output_compress = output_compress.view(batch_size, num_heads,
seq_len // block_elements, 1,
head_dim)
output_compress = output_compress.repeat(1, 1, 1, block_elements,
1).view(batch_size, num_heads,
seq_len, head_dim)
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num = generate_topk_block_sparse_pattern(
block_attn_score, topk)
output_select = block_sparse_attn(q, k, v, q2k_block_sparse_index,
q2k_block_sparse_num,
k2q_block_sparse_index,
k2q_block_sparse_num)
if compress_attn_weight is not None:
final_output = output_compress * compress_attn_weight + output_select
else:
final_output = output_compress + output_select
return final_output
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
QK = torch.matmul(q, k.transpose(-2, -1))
QK /= (q.size(-1)**0.5)
# Causal mask removed since causal is always false
QK = torch.nn.functional.softmax(QK, dim=-1)
output = torch.matmul(QK, v)
return output, QK
def generate_topk_block_sparse_pattern(block_attn_score: torch.Tensor,
topk: int):
"""
Generate a block sparse pattern where each q block attends to exactly topk kv blocks,
based on the provided attention scores.
Args:
block_attn_score: [bs, h, num_q_blocks, num_kv_blocks]
Attention scores between query and key blocks
topk: int
Number of kv blocks each q block attends to
Returns:
q2k_block_sparse_index: [bs, h, num_q_blocks, topk]
Contains the indices of kv blocks that each q block attends to.
q2k_block_sparse_num: [bs, h, num_q_blocks]
Contains the number of kv blocks that each q block attends to (all equal to topk).
k2q_block_sparse_index: [bs, h, num_kv_blocks, max_q_per_kv]
Contains the indices of q blocks that attend to each kv block.
k2q_block_sparse_num: [bs, h, num_kv_blocks]
Contains the number of q blocks that attend to each kv block.
"""
device = block_attn_score.device
# Extract dimensions from block_attn_score
bs, h, num_q_blocks, num_kv_blocks = block_attn_score.shape
sorted_result = torch.sort(block_attn_score, dim=-1, descending=True)
sorted_indice = sorted_result.indices
q2k_block_sparse_index, _ = torch.sort(sorted_indice[:, :, :, :topk],
dim=-1)
q2k_block_sparse_index = q2k_block_sparse_index.to(dtype=torch.int32)
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks),
topk,
device=device,
dtype=torch.int32)
block_map = topk_index_to_map(q2k_block_sparse_index,
num_kv_blocks,
transpose_map=True)
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(
block_map.transpose(2, 3))
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
@torch._dynamo.disable
def block_sparse_attn(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
@@ -61,6 +144,20 @@ def block_sparse_attn(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
)
def block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num):
"""
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks].
[*, *, i, j] = 1 means the i-th q block should attend to the j-th kv block.
"""
# assert all elements in q2k_block_sparse_num can be devisible by 2
o, lse = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
return o, lse
def block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num):
grad_output = grad_output.contiguous()
grad_q, grad_k, grad_v = block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
return grad_q, grad_k, grad_v
## pytorch sdpa version of block sparse ##
import triton
import triton.language as tl
@@ -128,6 +225,169 @@ def index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, BLOCK_Q, BLOCK_K
return mask
@triton.jit
def topk_index_to_map_kernel(
map_ptr,
index_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
topk: tl.constexpr,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
for i in tl.static_range(topk):
index = tl.load(index_ptr_base + i * index_kv_stride)
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
@triton.jit
def map_to_index_kernel(
map_ptr,
index_ptr,
index_num_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
index_num_bs_stride,
index_num_h_stride,
index_num_q_stride,
num_kv_blocks: tl.constexpr,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
num = 0
for i in tl.static_range(num_kv_blocks):
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
if map_entry:
tl.store(index_ptr_base + num * index_kv_stride, i)
num += 1
tl.store(
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
q * index_num_q_stride, num)
def topk_index_to_map(index: torch.Tensor,
num_kv_blocks: int,
transpose_map: bool = False):
"""
Convert topk indices to a map.
Args:
index: [bs, h, num_q_blocks, topk]
The topk indices tensor.
num_kv_blocks: int
The number of key-value blocks in the block_map returned
transpose_map: bool
If True, the block_map will be transposed on the final two dimensions.
Returns:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
A binary map where 1 indicates that the q block attends to the kv block.
"""
bs, h, num_q_blocks, topk = index.shape
if transpose_map is False:
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
dtype=torch.bool,
device=index.device)
else:
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
dtype=torch.bool,
device=index.device)
block_map = block_map.transpose(2, 3)
grid = (bs, h, num_q_blocks)
topk_index_to_map_kernel[grid](
block_map,
index,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
topk=topk,
)
return block_map
def map_to_index(block_map: torch.Tensor):
"""
Convert a block map to indices and counts.
Args:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
The block map tensor.
Returns:
index: [bs, h, num_q_blocks, num_kv_blocks]
The indices of the blocks.
index_num: [bs, h, num_q_blocks]
The number of blocks for each q block.
"""
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
index = torch.full((block_map.shape),
-1,
dtype=torch.int32,
device=block_map.device)
index_num = torch.empty((bs, h, num_q_blocks),
dtype=torch.int32,
device=block_map.device)
grid = (bs, h, num_q_blocks)
map_to_index_kernel[grid](
block_map,
index,
index_num,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
index_num.stride(0),
index_num.stride(1),
index_num.stride(2),
num_kv_blocks=num_kv_blocks,
)
return index, index_num
class BlockSparseAttentionFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
o, lse = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
ctx.save_for_backward(q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num)
return o
@staticmethod
def backward(ctx, grad_output):
q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num = ctx.saved_tensors
grad_q, grad_k, grad_v = block_sparse_attention_backward(
q, k, v, o, lse, grad_output, k2q_block_sparse_index, k2q_block_sparse_num
)
return grad_q, grad_k, grad_v, None, None, None, None
class DummyOperator(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
@@ -207,4 +467,4 @@ class BlockSparseAttnTorch:
o = DummyOperator.apply(output)
o.register_hook(self.recompute)
return o
return o
+23 -19
View File
@@ -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();
}
+44 -20
View File
@@ -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
+7 -1
View File
@@ -1,7 +1,7 @@
(sta-demo)=
# 🔍 Demo
There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
This is is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
<div style="text-align: center;">
<video controls width="800">
@@ -9,3 +9,9 @@ There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
Your browser does not support the video tag.
</video>
</div>
You can run STA using the following command:
```bash
bash scripts/inference/v1_inference_wan_STA.sh
```
+15 -45
View File
@@ -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.
+7
View File
@@ -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:
+1 -1
View File
@@ -10,7 +10,7 @@ def main():
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
# if num_gpus > 1, FastVideo will automatically handle distributed setup
# FastVideo will automatically handle distributed setup
num_gpus=2,
use_fsdp_inference=True,
use_cpu_offload=False
@@ -5,7 +5,7 @@ export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
base_port=29503
num_gpu=$(nvidia-smi --query-gpu=gpu_name --format=csv,noheader | wc -l)
num_gpu=1
gpu_ids=$(seq 0 $((num_gpu-1)))
skip_time_steps=12
@@ -14,7 +14,7 @@ STA_mode="STA_searching"
for i in $gpu_ids; do
port=$((base_port+i))
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
--prompt_path ./assets/prompt_extend_${i}.txt \
--prompt_path ./assets/prompt_${i}.txt \
--output_path $output_path \
--STA_mode $STA_mode &
sleep 1
@@ -27,7 +27,7 @@ STA_mode="STA_tuning"
for i in $gpu_ids; do
port=$((base_port+i))
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
--prompt_path ./assets/prompt_extend_${i}.txt \
--prompt_path ./assets/prompt_${i}.txt \
--output_path $output_path \
--STA_mode $STA_mode \
--skip_time_steps $skip_time_steps &
@@ -0,0 +1,10 @@
This directory contain e2e examples scripts for finetuning Wan2.1 I2V.
Execute the following commands from `FastVideo/` to run training:
- Download crush-smol dataset:
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/download_dataset.sh`
- Preprocess the videos and captions into latents:
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/preprocess_wan_data_i2v.sh`
- Edit the following file and run finetuning:
`bash examples/training/finetune/wan_i2v_14b_480p/crush_smol/finetune_i2v.sh`
@@ -0,0 +1,3 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -0,0 +1,91 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
DATA_DIR="data/crush-smol_processed_i2v/combined_parquet_dataset/"
VALIDATION_DIR="data/crush-smol_processed_i2v/validation_parquet_dataset/"
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_i2v_finetune"
--output_dir "$DATA_DIR/outputs/wan_i2v_finetune"
--max_train_steps 2000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 8
--num_height 480
--num_width 832
--num_frames 77
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 8
--tp_size 8
--hsdp_replicate_dim 1
--hsdp_shard_dim 8
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_preprocessed_path "$VALIDATION_DIR"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/v1/training/wan_i2v_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,130 @@
#!/bin/bash
#SBATCH --job-name=i2v
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=4
#SBATCH --ntasks=4
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
#SBATCH --mem=1440G
#SBATCH --output=i2v_output/i2v_%j.out
#SBATCH --error=i2v_output/i2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
DATA_DIR="data/crush-smol_processed_i2v/combined_parquet_dataset/"
VALIDATION_DIR="data/crush-smol_processed_i2v/validation_parquet_dataset/"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
# Training arguments
training_args=(
--tracker_project_name wan_i2v_finetune
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune_2n"
--max_train_steps=2000
--train_batch_size=2
--train_sp_batch_size 1
--gradient_accumulation_steps=1
--num_latent_t 8
--num_height 480
--num_width 832
--num_frames 77
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size $NUM_GPUS
--tp_size $NUM_GPUS
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES
--hsdp_shard_dim $NUM_GPUS
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 10
)
# Validation arguments
validation_args=(
--log_validation
--validation_preprocessed_path "$VALIDATION_DIR"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate=1e-5
--mixed_precision="bf16"
--checkpointing_steps=1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/v1/training/wan_i2v_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,25 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_i2v/"
VALIDATION_PATH="examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation.json"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--preprocess_task "i2v"
@@ -0,0 +1,31 @@
{
"data": [
{
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
"image_path": null,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -0,0 +1,10 @@
This directory contain e2e examples scripts for finetuning Wan2.1 T2v.
Execute the following commands from `FastVideo/` to run training:
- Download crush-smol dataset:
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/download_dataset.sh`
- Preprocess the videos and captions into latents:
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/preprocess_wan_data_t2v.sh`
- Edit the following file and run finetuning:
`bash examples/training/finetune/wan_t2v_1_3b/crush_smol/finetune_t2v.sh`
@@ -0,0 +1,3 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -0,0 +1,90 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DIR="data/crush-smol_processed_t2v/validation_parquet_dataset/"
NUM_GPUS=4
# export CUDA_VISIBLE_DEVICES=4,5
# Training arguments
training_args=(
--tracker_project_name "wan_t2v_finetune"
--output_dir "outputs/wan_t2v_finetune"
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 8
--num_latent_t 8
--num_height 480
--num_width 832
--num_frames 77
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size $NUM_GPUS
--tp_size $NUM_GPUS
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path $DATA_DIR
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_preprocessed_path $VALIDATION_DIR
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--checkpointing_steps 6000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/v1/training/wan_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,127 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
#SBATCH --mem=1440G
#SBATCH --output=t2v_output/t2v_%j.out
#SBATCH --error=t2v_output/t2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DIR="data/crush-smol_processed_t2v/validation_parquet_dataset/"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_finetune
--output_dir="outputs/wan_t2v_finetune"
--max_train_steps=1000
--train_batch_size=4
--train_sp_batch_size 1
--gradient_accumulation_steps=1
--num_latent_t 8
--num_height 480
--num_width 832
--num_frames 77
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 4
--tp_size 4
--hsdp_replicate_dim 2
--hsdp_shard_dim 4
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 10
)
# Validation arguments
validation_args=(
--log_validation
--validation_preprocessed_path "$VALIDATION_DIR"
--validation_steps 100
--validation_sampling_steps "50"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate=5e-5
--mixed_precision="bf16"
--checkpointing_steps=500
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/v1/training/wan_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,25 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
VALIDATION_PATH="examples/training/finetune/wan_t2v_1_3b/crush_smol/validation.json"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_dataset_file $VALIDATION_PATH \
--samples_per_file 8 \
--flush_frequency 8 \
--preprocess_task "t2v"
@@ -0,0 +1,31 @@
{
"data": [
{
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -104,13 +104,7 @@ if __name__ == "__main__":
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
args = parser.parse_args()
main(args)
-7
View File
@@ -671,13 +671,6 @@ if __name__ == "__main__":
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
-7
View File
@@ -693,13 +693,6 @@ if __name__ == "__main__":
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
+1 -7
View File
@@ -520,13 +520,7 @@ if __name__ == "__main__":
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
+2 -15
View File
@@ -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,
+1 -1
View File
@@ -10,8 +10,8 @@ from fastvideo.v1.attention.selector import get_attn_backend
__all__ = [
"DistributedAttention",
"DistributedAttention_VSA",
"LocalAttention",
"DistributedAttention_VSA",
"AttentionBackend",
"AttentionMetadata",
"AttentionMetadataBuilder",
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import json
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Type
from typing import Any, List, Optional, Type
import torch
from einops import rearrange
@@ -17,33 +17,11 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.utils import dict_to_3d_list
logger = init_logger(__name__)
# TODO(will-refactor): move this to a utils file
def dict_to_3d_list(
mask_strategy: Dict[str,
Any]) -> List[List[List[Optional[torch.Tensor]]]]:
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
max_timesteps_idx = max(
timesteps_idx for timesteps_idx, layer_idx, head_idx in indices) + 1
max_layer_idx = max(layer_idx
for timesteps_idx, layer_idx, head_idx in indices) + 1
max_head_idx = max(head_idx
for timesteps_idx, layer_idx, head_idx in indices) + 1
result = [[[None for _ in range(max_head_idx)]
for _ in range(max_layer_idx)] for _ in range(max_timesteps_idx)]
for key, value in mask_strategy.items():
timesteps_idx, layer_idx, head_idx = map(int, key.split('_'))
result[timesteps_idx][layer_idx][head_idx] = value
return result
class RangeDict(dict):
def __getitem__(self, item: int) -> str:
@@ -287,16 +265,10 @@ class SlidingTileAttentionImpl(AttentionImpl):
forward_batch.mask_search_final_result_pos[timestep].append(
layer_loss_save)
else:
# windows = [
# self.mask_strategy[timestep][layer_idx][head_idx + start_head]
# for head_idx in range(head_num)
# ]
windows = [
STA_param[head_idx + start_head] for head_idx in range(head_num)
]
# if has_text is False:
# from IPython import embed
# embed()
hidden_states = sliding_tile_attention(
query, key, value, windows, text_length, has_text,
self.dit_seq_shape_str).transpose(1, 2)
@@ -1,18 +1,15 @@
# SPDX-License-Identifier: Apache-2.0
import math
from dataclasses import dataclass
from typing import Any, List, Optional, Type, cast
from typing import List, Optional, Type
import torch
import triton
import triton.language as tl
from einops import rearrange
try:
from vsa import block_sparse_attn
except ImportError: # noqa: E722
block_sparse_attn = None
from typing import Tuple
from vsa import video_sparse_attn
except ImportError:
video_sparse_attn = None
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
@@ -75,14 +72,18 @@ class VideoSparseAttentionMetadataBuilder(AttentionMetadataBuilder):
if forward_batch.latents is None:
raise ValueError("latents cannot be None")
raw_latent_shape = forward_batch.latents.shape
patch_size = fastvideo_args.dit_config.patch_size
raw_latent_shape = forward_batch.raw_latent_shape
if raw_latent_shape is None:
raise ValueError("raw_latent_shape cannot be None")
patch_size = fastvideo_args.pipeline_config.dit_config.patch_size
dit_seq_shape = [
raw_latent_shape[2] // patch_size[0],
raw_latent_shape[3] // patch_size[1],
raw_latent_shape[4] // patch_size[2]
]
VSA_sparsity = forward_batch.VSA_sparsity
return VideoSparseAttentionMetadata(current_timestep=current_timestep,
dit_seq_shape=dit_seq_shape,
VSA_sparsity=VSA_sparsity)
@@ -177,12 +178,16 @@ class VideoSparseAttentionImpl(AttentionImpl):
value = value.transpose(1, 2).contiguous()
gate_compress = gate_compress.transpose(1, 2).contiguous()
VSA_sparsity = attn_metadata.VSA_sparsity
cur_topk = math.ceil(
(1 - attn_metadata.VSA_sparsity) *
(1 - VSA_sparsity) *
(self.img_seq_length / math.prod(self.VSA_base_tile_size)))
# Cast to Any to bypass type checking for untyped function
hidden_states = cast(Any, sparse_attn_c_s_p)(
if video_sparse_attn is None:
raise NotImplementedError("video_sparse_attn is not installed")
hidden_states = video_sparse_attn(
query,
key,
value,
@@ -191,267 +196,3 @@ class VideoSparseAttentionImpl(AttentionImpl):
compress_attn_weight=gate_compress).transpose(1, 2)
return hidden_states
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
QK = torch.matmul(q, k.transpose(-2, -1))
QK /= (q.size(-1)**0.5)
# Causal mask removed since causal is always false
QK = torch.nn.functional.softmax(QK, dim=-1)
output = torch.matmul(QK, v)
return output, QK
def sparse_attn_c_s_p(q, k, v, topk, block_size, compress_attn_weight=None):
"""
q: [batch_size, num_heads, seq_len, head_dim]
k: [batch_size, num_heads, seq_len, head_dim]
v: [batch_size, num_heads, seq_len, head_dim]
topk: int
block_size: int or tuple of 3 ints
video_shape: tuple of (T, H, W)
compress_attn_weight: [batch_size, num_heads, seq_len, head_dim]
select_attn_weight: [batch_size, num_heads, seq_len, head_dim]
V1 of sparse attention. Include compress attn and sparse attn branch, use average pooling to compress.
Assume q, k, v is flattened in this way: [batch_size, num_heads, T//block_size[0], H//block_size[1], W//block_size[2], block_size[0], block_size[1], block_size[2]]
"""
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
block_elements = block_size[0] * block_size[1] * block_size[2]
assert block_elements % 64 == 0 and block_elements >= 64
assert q.shape[2] % block_elements == 0
batch_size, num_heads, seq_len, head_dim = q.shape
# compress attn
q_compress = q.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
k_compress = k.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
v_compress = v.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
v_compress)
output_compress = output_compress.view(batch_size, num_heads,
seq_len // block_elements, 1,
head_dim)
output_compress = output_compress.repeat(1, 1, 1, block_elements,
1).view(batch_size, num_heads,
seq_len, head_dim)
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num = generate_topk_block_sparse_pattern(
block_attn_score, topk)
output_select = block_sparse_attn(q, k, v, q2k_block_sparse_index,
q2k_block_sparse_num,
k2q_block_sparse_index,
k2q_block_sparse_num)
if compress_attn_weight is not None:
final_output = output_compress * compress_attn_weight + output_select
else:
final_output = output_compress + output_select
return final_output
@triton.jit
def topk_index_to_map_kernel(
map_ptr,
index_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
topk: tl.constexpr,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
for i in tl.static_range(topk):
index = tl.load(index_ptr_base + i * index_kv_stride)
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
@triton.jit
def map_to_index_kernel(
map_ptr,
index_ptr,
index_num_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
index_num_bs_stride,
index_num_h_stride,
index_num_q_stride,
num_kv_blocks: tl.constexpr,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
num = 0
for i in tl.static_range(num_kv_blocks):
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
if map_entry:
tl.store(index_ptr_base + num * index_kv_stride, i)
num += 1
tl.store(
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
q * index_num_q_stride, num)
def topk_index_to_map(index: torch.Tensor,
num_kv_blocks: int,
transpose_map: bool = False):
"""
Convert topk indices to a map.
Args:
index: [bs, h, num_q_blocks, topk]
The topk indices tensor.
num_kv_blocks: int
The number of key-value blocks in the block_map returned
transpose_map: bool
If True, the block_map will be transposed on the final two dimensions.
Returns:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
A binary map where 1 indicates that the q block attends to the kv block.
"""
bs, h, num_q_blocks, topk = index.shape
if transpose_map is False:
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
dtype=torch.bool,
device=index.device)
else:
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
dtype=torch.bool,
device=index.device)
block_map = block_map.transpose(2, 3)
grid = (bs, h, num_q_blocks)
topk_index_to_map_kernel[grid](
block_map,
index,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
topk=topk,
)
return block_map
def map_to_index(block_map: torch.Tensor):
"""
Convert a block map to indices and counts.
Args:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
The block map tensor.
Returns:
index: [bs, h, num_q_blocks, num_kv_blocks]
The indices of the blocks.
index_num: [bs, h, num_q_blocks]
The number of blocks for each q block.
"""
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
index = torch.full((block_map.shape),
-1,
dtype=torch.int32,
device=block_map.device)
index_num = torch.empty((bs, h, num_q_blocks),
dtype=torch.int32,
device=block_map.device)
grid = (bs, h, num_q_blocks)
map_to_index_kernel[grid](
block_map,
index,
index_num,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
index_num.stride(0),
index_num.stride(1),
index_num.stride(2),
num_kv_blocks=num_kv_blocks,
)
return index, index_num
def generate_topk_block_sparse_pattern(block_attn_score: torch.Tensor,
topk: int):
"""
Generate a block sparse pattern where each q block attends to exactly topk kv blocks,
based on the provided attention scores.
Args:
block_attn_score: [bs, h, num_q_blocks, num_kv_blocks]
Attention scores between query and key blocks
topk: int
Number of kv blocks each q block attends to
Returns:
q2k_block_sparse_index: [bs, h, num_q_blocks, topk]
Contains the indices of kv blocks that each q block attends to.
q2k_block_sparse_num: [bs, h, num_q_blocks]
Contains the number of kv blocks that each q block attends to (all equal to topk).
k2q_block_sparse_index: [bs, h, num_kv_blocks, max_q_per_kv]
Contains the indices of q blocks that attend to each kv block.
k2q_block_sparse_num: [bs, h, num_kv_blocks]
Contains the number of q blocks that attend to each kv block.
"""
device = block_attn_score.device
# Extract dimensions from block_attn_score
bs, h, num_q_blocks, num_kv_blocks = block_attn_score.shape
sorted_result = torch.sort(block_attn_score, dim=-1, descending=True)
sorted_indice = sorted_result.indices
q2k_block_sparse_index, _ = torch.sort(sorted_indice[:, :, :, :topk],
dim=-1)
q2k_block_sparse_index = q2k_block_sparse_index.to(dtype=torch.int32)
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks),
topk,
device=device,
dtype=torch.int32)
block_map = topk_index_to_map(q2k_block_sparse_index,
num_kv_blocks,
transpose_map=True)
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(
block_map.transpose(2, 3))
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
+26 -27
View File
@@ -12,7 +12,7 @@ from fastvideo.v1.distributed.communication_op import (
from fastvideo.v1.distributed.parallel_state import (get_sp_parallel_rank,
get_sp_world_size)
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
from fastvideo.v1.utils import get_compute_dtype
@@ -26,8 +26,8 @@ class DistributedAttention(nn.Module):
num_kv_heads: Optional[int] = None,
softmax_scale: Optional[float] = None,
causal: bool = False,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
supported_attention_backends: Optional[Tuple[
AttentionBackendEnum, ...]] = None,
prefix: str = "",
**extra_impl_args) -> None:
super().__init__()
@@ -45,13 +45,13 @@ class DistributedAttention(nn.Module):
dtype,
supported_attention_backends=supported_attention_backends)
impl_cls = attn_backend.get_impl_cls()
self.impl = impl_cls(num_heads=num_heads,
head_size=head_size,
causal=causal,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
prefix=f"{prefix}.impl",
**extra_impl_args)
self.attn_impl = impl_cls(num_heads=num_heads,
head_size=head_size,
causal=causal,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
prefix=f"{prefix}.impl",
**extra_impl_args)
self.num_heads = num_heads
self.head_size = head_size
self.num_kv_heads = num_kv_heads
@@ -100,7 +100,7 @@ class DistributedAttention(nn.Module):
scatter_dim=2,
gather_dim=1)
# Apply backend-specific preprocess_qkv
qkv = self.impl.preprocess_qkv(qkv, ctx_attn_metadata)
qkv = self.attn_impl.preprocess_qkv(qkv, ctx_attn_metadata)
# Concatenate with replicated QKV if provided
if replicated_q is not None:
@@ -116,7 +116,7 @@ class DistributedAttention(nn.Module):
q, k, v = qkv.chunk(3, dim=0)
output = self.impl.forward(q, k, v, ctx_attn_metadata)
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
# Redistribute back if using sequence parallelism
replicated_output = None
@@ -127,7 +127,7 @@ class DistributedAttention(nn.Module):
replicated_output = sequence_model_parallel_all_gather(
replicated_output.contiguous(), dim=2)
# Apply backend-specific postprocess_output
output = self.impl.postprocess_output(output, ctx_attn_metadata)
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
output = sequence_model_parallel_all_to_all_4D(output,
scatter_dim=1,
@@ -183,18 +183,17 @@ class DistributedAttention_VSA(DistributedAttention):
scatter_dim=2,
gather_dim=1)
qkvg = self.impl.preprocess_qkv(
qkvg, ctx_attn_metadata) # (yongqi) pass latent shape here?
qkvg = self.attn_impl.preprocess_qkv(qkvg, ctx_attn_metadata)
q, k, v, gate_compress = qkvg.chunk(4, dim=0)
output = self.impl.forward(q, k, v, gate_compress,
ctx_attn_metadata) # type: ignore[call-arg]
output = self.attn_impl.forward(
q, k, v, gate_compress, ctx_attn_metadata) # type: ignore[call-arg]
# Redistribute back if using sequence parallelism
replicated_output = None
# Apply backend-specific postprocess_output
output = self.impl.postprocess_output(output, ctx_attn_metadata)
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
output = sequence_model_parallel_all_to_all_4D(output,
scatter_dim=1,
@@ -212,8 +211,8 @@ class LocalAttention(nn.Module):
num_kv_heads: Optional[int] = None,
softmax_scale: Optional[float] = None,
causal: bool = False,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
supported_attention_backends: Optional[Tuple[
AttentionBackendEnum, ...]] = None,
**extra_impl_args) -> None:
super().__init__()
if softmax_scale is None:
@@ -229,12 +228,12 @@ class LocalAttention(nn.Module):
dtype,
supported_attention_backends=supported_attention_backends)
impl_cls = attn_backend.get_impl_cls()
self.impl = impl_cls(num_heads=num_heads,
head_size=head_size,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
causal=causal,
**extra_impl_args)
self.attn_impl = impl_cls(num_heads=num_heads,
head_size=head_size,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
causal=causal,
**extra_impl_args)
self.num_heads = num_heads
self.head_size = head_size
self.num_kv_heads = num_kv_heads
@@ -265,5 +264,5 @@ class LocalAttention(nn.Module):
forward_context: ForwardContext = get_forward_context()
ctx_attn_metadata = forward_context.attn_metadata
output = self.impl.forward(q, k, v, ctx_attn_metadata)
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
return output
+14 -11
View File
@@ -11,13 +11,13 @@ import torch
import fastvideo.v1.envs as envs
from fastvideo.v1.attention.backends.abstract import AttentionBackend
from fastvideo.v1.logger import init_logger
from fastvideo.v1.platforms import _Backend, current_platform
from fastvideo.v1.platforms import AttentionBackendEnum, current_platform
from fastvideo.v1.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
logger = init_logger(__name__)
def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
def backend_name_to_enum(backend_name: str) -> Optional[AttentionBackendEnum]:
"""
Convert a string backend name to a _Backend enum value.
@@ -27,11 +27,11 @@ def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
loaded.
"""
assert backend_name is not None
return _Backend[backend_name] if backend_name in _Backend.__members__ else \
return AttentionBackendEnum[backend_name] if backend_name in AttentionBackendEnum.__members__ else \
None
def get_env_variable_attn_backend() -> Optional[_Backend]:
def get_env_variable_attn_backend() -> Optional[AttentionBackendEnum]:
'''
Get the backend override specified by the FastVideo attention
backend environment variable, if one is specified.
@@ -53,10 +53,11 @@ def get_env_variable_attn_backend() -> Optional[_Backend]:
#
# THIS SELECTION TAKES PRECEDENCE OVER THE
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
forced_attn_backend: Optional[_Backend] = None
forced_attn_backend: Optional[AttentionBackendEnum] = None
def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
def global_force_attn_backend(
attn_backend: Optional[AttentionBackendEnum]) -> None:
'''
Force all attention operations to use a specified backend.
@@ -71,7 +72,7 @@ def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
forced_attn_backend = attn_backend
def get_global_forced_attn_backend() -> Optional[_Backend]:
def get_global_forced_attn_backend() -> Optional[AttentionBackendEnum]:
'''
Get the currently-forced choice of attention backend,
or None if auto-selection is currently enabled.
@@ -82,7 +83,8 @@ def get_global_forced_attn_backend() -> Optional[_Backend]:
def get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
) -> Type[AttentionBackend]:
return _cached_get_attn_backend(head_size, dtype,
supported_attention_backends)
@@ -92,7 +94,8 @@ def get_attn_backend(
def _cached_get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
) -> Type[AttentionBackend]:
# Check whether a particular choice of backend was
# previously forced.
@@ -102,7 +105,7 @@ def _cached_get_attn_backend(
if not supported_attention_backends:
raise ValueError("supported_attention_backends is empty")
selected_backend = None
backend_by_global_setting: Optional[_Backend] = (
backend_by_global_setting: Optional[AttentionBackendEnum] = (
get_global_forced_attn_backend())
if backend_by_global_setting is not None:
selected_backend = backend_by_global_setting
@@ -125,7 +128,7 @@ def _cached_get_attn_backend(
@contextmanager
def global_force_attn_backend_context_manager(
attn_backend: _Backend) -> Generator[None, None, None]:
attn_backend: AttentionBackendEnum) -> Generator[None, None, None]:
'''
Globally force a FastVideo attention backend override within a
context manager, reverting the global attention backend
+4 -2
View File
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field, fields
from typing import Any, Dict
from typing import Any, Dict, List, Tuple
from fastvideo.v1.logger import init_logger
@@ -12,7 +12,9 @@ logger = init_logger(__name__)
# 3. Any field in ArchConfig is fixed upon initialization, and should be hidden away from users
@dataclass
class ArchConfig:
pass
stacked_params_mapping: List[Tuple[str, str, str]] = field(
default_factory=list
) # mapping from huggingface weight names to custom names
@dataclass
+6 -7
View File
@@ -4,7 +4,7 @@ from typing import Any, List, Optional, Tuple
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.layers.quantization import QuantizationConfig
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
@dataclass
@@ -12,13 +12,12 @@ class DiTArchConfig(ArchConfig):
_fsdp_shard_conditions: list = field(default_factory=list)
_compile_conditions: list = field(default_factory=list)
_param_names_mapping: dict = field(default_factory=dict)
_reverse_param_names_mapping: dict = field(default_factory=dict)
_lora_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: Tuple[_Backend,
...] = (_Backend.SLIDING_TILE_ATTN,
_Backend.SAGE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA,
_Backend.VIDEO_SPARSE_ATTN)
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN)
hidden_size: int = 0
num_attention_heads: int = 0
@@ -147,6 +147,9 @@ class HunyuanVideoArchConfig(DiTArchConfig):
r"final_layer.linear.\1",
})
# Reverse mapping for saving checkpoints: training -> diffusers
_reverse_param_names_mapping: dict = field(default_factory=lambda: {})
patch_size: int = 2
patch_size_t: int = 1
in_channels: int = 16
@@ -5,13 +5,11 @@ from typing import List, Optional, Tuple, Union
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class StepVideoArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
_fsdp_shard_conditions: list = field(
default_factory=lambda:
[lambda n, m: "transformer_blocks" in n and n.split(".")[-1].isdigit()])
_param_names_mapping: dict = field(
default_factory=lambda: {
+5 -1
View File
@@ -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(
+7 -4
View File
@@ -6,14 +6,14 @@ import torch
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.layers.quantization import QuantizationConfig
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
@dataclass
class EncoderArchConfig(ArchConfig):
architectures: List[str] = field(default_factory=lambda: [])
_supported_attention_backends: Tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA)
output_hidden_states: bool = False
use_return_dict: bool = True
@@ -32,8 +32,11 @@ class TextEncoderArchConfig(EncoderArchConfig):
output_past: bool = True
scalable_attention: bool = True
tie_word_embeddings: bool = False
stacked_params_mapping: List[Tuple[str, str, str]] = field(
default_factory=list
) # mapping from huggingface weight names to custom names
tokenizer_kwargs: Dict[str, Any] = field(default_factory=dict)
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
def __post_init__(self) -> None:
self.tokenizer_kwargs = {
+18 -1
View File
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import Optional
from typing import List, Optional, Tuple
from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
ImageEncoderConfig,
@@ -8,6 +8,14 @@ from fastvideo.v1.configs.models.encoders.base import (ImageEncoderArchConfig,
TextEncoderConfig)
def _is_transformer_layer(n: str, m) -> bool:
return "layers" in n and str.isdigit(n.split(".")[-1])
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embeddings")
@dataclass
class CLIPTextArchConfig(TextEncoderArchConfig):
vocab_size: int = 49408
@@ -27,6 +35,15 @@ class CLIPTextArchConfig(TextEncoderArchConfig):
bos_token_id: int = 49406
eos_token_id: int = 49407
text_len: int = 77
stacked_params_mapping: List[Tuple[str, str,
str]] = field(default_factory=lambda: [
# (param_name, shard_name, shard_id)
("qkv_proj", "q_proj", "q"),
("qkv_proj", "k_proj", "k"),
("qkv_proj", "v_proj", "v"),
])
_fsdp_shard_conditions: list = field(
default_factory=lambda: [_is_transformer_layer, _is_embeddings])
@dataclass
+25 -1
View File
@@ -1,11 +1,23 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import Optional
from typing import List, Optional, Tuple
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig)
def _is_transformer_layer(n: str, m) -> bool:
return "layers" in n and str.isdigit(n.split(".")[-1])
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embed_tokens")
def _is_final_norm(n: str, m) -> bool:
return n.endswith("norm")
@dataclass
class LlamaArchConfig(TextEncoderArchConfig):
vocab_size: int = 32000
@@ -32,6 +44,18 @@ class LlamaArchConfig(TextEncoderArchConfig):
head_dim: Optional[int] = None
hidden_state_skip_layer: int = 2
text_len: int = 256
stacked_params_mapping: List[Tuple[str, str, str]] = field(
default_factory=lambda: [
# (param_name, shard_name, shard_id)
(".qkv_proj", ".q_proj", "q"),
(".qkv_proj", ".k_proj", "k"),
(".qkv_proj", ".v_proj", "v"),
(".gate_up_proj", ".gate_proj", 0), # type: ignore
(".gate_up_proj", ".up_proj", 1), # type: ignore
])
_fsdp_shard_conditions: list = field(
default_factory=lambda:
[_is_transformer_layer, _is_embeddings, _is_final_norm])
@dataclass
+23 -1
View File
@@ -1,11 +1,23 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import Optional
from typing import List, Optional, Tuple
from fastvideo.v1.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig)
def _is_transformer_layer(n: str, m) -> bool:
return "block" in n and str.isdigit(n.split(".")[-1])
def _is_embeddings(n: str, m) -> bool:
return n.endswith("shared")
def _is_final_layernorm(n: str, m) -> bool:
return n.endswith("final_layer_norm")
@dataclass
class T5ArchConfig(TextEncoderArchConfig):
vocab_size: int = 32128
@@ -29,6 +41,16 @@ class T5ArchConfig(TextEncoderArchConfig):
eos_token_id: int = 1
classifier_dropout: float = 0.0
text_len: int = 512
stacked_params_mapping: List[Tuple[str, str,
str]] = field(default_factory=lambda: [
# (param_name, shard_name, shard_id)
(".qkv_proj", ".q", "q"),
(".qkv_proj", ".k", "k"),
(".qkv_proj", ".v", "v"),
])
_fsdp_shard_conditions: list = field(
default_factory=lambda:
[_is_transformer_layer, _is_embeddings, _is_final_layernorm])
# Referenced from https://github.com/huggingface/transformers/blob/main/src/transformers/models/t5/configuration_t5.py
def __post_init__(self):
+11
View File
@@ -1,4 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import dataclasses
from dataclasses import dataclass, field
from typing import Any, Union
@@ -129,3 +131,12 @@ class VAEConfig(ModelConfig):
)
return parser
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "VAEConfig":
kwargs = {}
for attr in dataclasses.fields(cls):
value = getattr(args, attr.name, None)
if value is not None:
kwargs[attr.name] = value
return cls(**kwargs)
+2 -2
View File
@@ -3,7 +3,7 @@ from fastvideo.v1.configs.pipelines.base import (PipelineConfig,
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
HunyuanConfig)
from fastvideo.v1.configs.pipelines.registry import (
get_pipeline_config_cls_for_name)
get_pipeline_config_cls_from_name)
from fastvideo.v1.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.v1.configs.pipelines.wan import (WanI2V480PConfig,
WanI2V720PConfig,
@@ -14,5 +14,5 @@ __all__ = [
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
"get_pipeline_config_cls_for_name"
"get_pipeline_config_cls_from_name"
]
+253 -18
View File
@@ -1,19 +1,31 @@
# SPDX-License-Identifier: Apache-2.0
import json
from dataclasses import asdict, dataclass, field, fields
from typing import Any, Callable, Dict, Optional, Tuple, cast
from enum import Enum
from typing import Any, Callable, Dict, List, Optional, Tuple, Union, cast
import torch
from fastvideo.v1.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
VAEConfig)
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput
from fastvideo.v1.configs.utils import update_config_from_args
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import shallow_asdict
from fastvideo.v1.utils import (FlexibleArgumentParser, StoreBoolean,
shallow_asdict)
logger = init_logger(__name__)
class STA_Mode(str, Enum):
"""STA (Sliding Tile Attention) modes."""
STA_INFERENCE = "STA_inference"
STA_SEARCHING = "STA_searching"
STA_TUNING = "STA_tuning"
STA_TUNING_CFG = "STA_tuning_cfg"
NONE = None
def preprocess_text(prompt: str) -> str:
return prompt
@@ -22,59 +34,282 @@ def postprocess_text(output: BaseEncoderOutput) -> torch.tensor:
raise NotImplementedError
# config for a single pipeline
@dataclass
class PipelineConfig:
"""Base configuration for all pipeline architectures."""
model_path: str = ""
pipeline_config_path: Optional[str] = None
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
disable_autocast: bool = False
# Model configuration
precision: str = "bf16"
dit_config: DiTConfig = field(default_factory=DiTConfig)
dit_precision: str = "bf16"
# VAE configuration
vae_config: VAEConfig = field(default_factory=VAEConfig)
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = True
vae_config: VAEConfig = field(default_factory=VAEConfig)
# DiT configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
# Image encoder configuration
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
image_encoder_precision: str = "fp32"
# Text encoder configuration
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp16", ))
DEFAULT_TEXT_ENCODER_PRECISIONS = ("fp16", )
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (EncoderConfig(), ))
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp16", ))
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
...] = field(default_factory=lambda:
(postprocess_text, ))
# STA (Spatial-Temporal Attention) parameters
# LoRA parameters
lora_path: Optional[str] = None
lora_nickname: Optional[
str] = "default" # for swapping adapters in the pipeline
lora_target_names: Optional[List[
str]] = None # can restrict list of layers to adapt, e.g. ["q_proj"]
# StepVideo specific parameters
pos_magic: Optional[str] = None
neg_magic: Optional[str] = None
timesteps_scale: Optional[bool] = None
# STA (Sliding Tile Attention) parameters
mask_strategy_file_path: Optional[str] = None
STA_mode: str = "STA_inference"
STA_mode: STA_Mode = STA_Mode.STA_INFERENCE
skip_time_steps: int = 15
# Compilation
enable_torch_compile: bool = False
# enable_torch_compile: bool = False
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser,
prefix: str = "") -> FlexibleArgumentParser:
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
# model_path will be conflicting with the model_path in FastVideoArgs,
# so we add it separately if prefix is not empty
if prefix_with_dot != "":
parser.add_argument(
f"--{prefix_with_dot}model-path",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}model_path",
default=PipelineConfig.model_path,
help="Path to the pretrained model",
)
parser.add_argument(
f"--{prefix_with_dot}pipeline-config-path",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}pipeline_config_path",
default=PipelineConfig.pipeline_config_path,
help="Path to the pipeline config",
)
parser.add_argument(
f"--{prefix_with_dot}embedded-cfg-scale",
type=float,
dest=f"{prefix_with_dot.replace('-', '_')}embedded_cfg_scale",
default=PipelineConfig.embedded_cfg_scale,
help="Embedded CFG scale",
)
parser.add_argument(
f"--{prefix_with_dot}flow-shift",
type=float,
dest=f"{prefix_with_dot.replace('-', '_')}flow_shift",
default=PipelineConfig.flow_shift,
help="Flow shift parameter",
)
# DiT configuration
parser.add_argument(
f"--{prefix_with_dot}dit-precision",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}dit_precision",
default=PipelineConfig.dit_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for the DiT model",
)
# VAE configuration
parser.add_argument(
f"--{prefix_with_dot}vae-precision",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}vae_precision",
default=PipelineConfig.vae_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for VAE",
)
parser.add_argument(
f"--{prefix_with_dot}vae-tiling",
action=StoreBoolean,
dest=f"{prefix_with_dot.replace('-', '_')}vae_tiling",
default=PipelineConfig.vae_tiling,
help="Enable VAE tiling",
)
parser.add_argument(
f"--{prefix_with_dot}vae-sp",
action=StoreBoolean,
dest=f"{prefix_with_dot.replace('-', '_')}vae_sp",
help="Enable VAE spatial parallelism",
)
# Text encoder configuration
parser.add_argument(
f"--{prefix_with_dot}text-encoder-precisions",
nargs="+",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}text_encoder_precisions",
default=PipelineConfig.DEFAULT_TEXT_ENCODER_PRECISIONS,
choices=["fp32", "fp16", "bf16"],
help="Precision for each text encoder",
)
# Image encoder configuration
parser.add_argument(
f"--{prefix_with_dot}image-encoder-precision",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}image_encoder_precision",
default=PipelineConfig.image_encoder_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for image encoder",
)
parser.add_argument(
f"--{prefix_with_dot}pos_magic",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}pos_magic",
default=PipelineConfig.pos_magic,
help="Positive magic prompt for sampling, used in stepvideo",
)
parser.add_argument(
f"--{prefix_with_dot}neg_magic",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}neg_magic",
default=PipelineConfig.neg_magic,
help="Negative magic prompt for sampling, used in stepvideo",
)
parser.add_argument(
f"--{prefix_with_dot}timesteps_scale",
type=bool,
dest=f"{prefix_with_dot.replace('-', '_')}timesteps_scale",
default=PipelineConfig.timesteps_scale,
help=
"Bool for applying scheduler scale in set_timesteps, used in stepvideo",
)
# Add VAE configuration arguments
from fastvideo.v1.configs.models.vaes.base import VAEConfig
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
# Add DiT configuration arguments
from fastvideo.v1.configs.models.dits.base import DiTConfig
DiTConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}dit-config")
return parser
def update_config_from_dict(self,
args: Dict[str, Any],
prefix: str = "") -> None:
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
update_config_from_args(self, args, prefix, pop_args=True)
update_config_from_args(self.vae_config,
args,
f"{prefix_with_dot}vae_config",
pop_args=True)
update_config_from_args(self.dit_config,
args,
f"{prefix_with_dot}dit_config",
pop_args=True)
@classmethod
def from_pretrained(cls, model_path: str) -> "PipelineConfig":
"""
use the pipeline class setting from model_path to match the pipeline config
"""
from fastvideo.v1.configs.pipelines.registry import (
get_pipeline_config_cls_for_name)
pipeline_config_cls = get_pipeline_config_cls_for_name(model_path)
if pipeline_config_cls is not None:
pipeline_config = pipeline_config_cls()
else:
get_pipeline_config_cls_from_name)
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
return cast(PipelineConfig, pipeline_config_cls(model_path=model_path))
@classmethod
def from_kwargs(cls,
kwargs: Dict[str, Any],
config_cli_prefix: str = "") -> "PipelineConfig":
"""
Load PipelineConfig from kwargs Dictionary.
kwargs: dictionary of kwargs
config_cli_prefix: prefix of CLI arguments for this PipelineConfig instance
"""
from fastvideo.v1.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
prefix_with_dot = f"{config_cli_prefix}." if (config_cli_prefix.strip()
!= "") else ""
model_path: Optional[str] = kwargs.get(prefix_with_dot + 'model_path',
None) or kwargs.get('model_path')
pipeline_config_or_path: Optional[Union[str, PipelineConfig, Dict[
str, Any]]] = kwargs.get(prefix_with_dot + 'pipeline_config',
None) or kwargs.get('pipeline_config')
if model_path is None:
raise ValueError("model_path is required in kwargs")
# 1. Get the pipeline config class from the registry
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
# 2. Instantiate PipelineConfig
if pipeline_config_cls is None:
logger.warning(
"Couldn't find an optimal sampling param for %s. Using the default sampling param.",
"Couldn't find pipeline config for %s. Using the default pipeline config.",
model_path)
pipeline_config = cls()
else:
pipeline_config = pipeline_config_cls()
return cast(PipelineConfig, pipeline_config)
# 3. Load PipelineConfig from a json file or a PipelineConfig object if provided
if isinstance(pipeline_config_or_path, str):
pipeline_config.load_from_json(pipeline_config_or_path)
kwargs[prefix_with_dot +
'pipeline_config_path'] = pipeline_config_or_path
elif isinstance(pipeline_config_or_path, PipelineConfig):
pipeline_config = pipeline_config_or_path
elif isinstance(pipeline_config_or_path, dict):
pipeline_config.update_pipeline_config(pipeline_config_or_path)
# 4. Update PipelineConfig from CLI arguments if provided
kwargs[prefix_with_dot + 'model_path'] = model_path
pipeline_config.update_config_from_dict(kwargs, config_cli_prefix)
return pipeline_config
def check_pipeline_config(self) -> None:
if self.vae_sp and not self.vae_tiling:
raise ValueError(
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
)
if len(self.text_encoder_configs) != len(self.text_encoder_precisions):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
)
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
if len(self.preprocess_text_funcs) != len(self.postprocess_text_funcs):
raise ValueError(
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
def dump_to_json(self, file_path: str):
output_dict = shallow_asdict(self)
+1 -1
View File
@@ -80,7 +80,7 @@ class HunyuanConfig(PipelineConfig):
(llama_postprocess_text, clip_postprocess_text))
# Precision for each component
precision: str = "bf16"
dit_precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp16", "fp16"))
+63 -26
View File
@@ -19,7 +19,7 @@ from fastvideo.v1.utils import (maybe_download_model_index,
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[PipelineConfig]] = {
PIPE_NAME_TO_CONFIG: Dict[str, Type[PipelineConfig]] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
@@ -51,37 +51,74 @@ PIPELINE_FALLBACK_CONFIG: Dict[str, Type[PipelineConfig]] = {
}
def get_pipeline_config_cls_for_name(
pipeline_name_or_path: str) -> Optional[type[PipelineConfig]]:
"""Get the appropriate config class for specific pretrained weights."""
def get_pipeline_config_cls_from_name(
pipeline_name_or_path: str) -> Type[PipelineConfig]:
"""Get the appropriate configuration class for a given pipeline name or path.
if os.path.exists(pipeline_name_or_path):
config = verify_model_config_and_directory(pipeline_name_or_path)
logger.warning(
"FastVideo may not correctly identify the optimal config for this model, as the local directory may have been renamed."
)
else:
config = maybe_download_model_index(pipeline_name_or_path)
This function implements a multi-step lookup process to find the most suitable
configuration class for a given pipeline. It follows this order:
1. Exact match in the PIPE_NAME_TO_CONFIG
2. Partial match in the PIPE_NAME_TO_CONFIG
3. Fallback to class name in the model_index.json
4. else raise an error
pipeline_name = config["_class_name"]
Args:
pipeline_name_or_path (str): The name or path of the pipeline. This can be:
- A registered model ID (e.g., "FastVideo/FastHunyuan-diffusers")
- A local path to a model directory
- A model ID that will be downloaded
Returns:
Type[PipelineConfig]: The configuration class that best matches the pipeline.
This will be one of:
- A specific weight configuration class if an exact match is found
- A fallback configuration class based on the pipeline architecture
- The base PipelineConfig class if no matches are found
Note:
- For local paths, the function will verify the model configuration
- For remote models, it will attempt to download the model index
- Warning messages are logged when falling back to less specific configurations
"""
pipeline_config_cls: Optional[Type[PipelineConfig]] = None
# First try exact match for specific weights
if pipeline_name_or_path in WEIGHT_CONFIG_REGISTRY:
return WEIGHT_CONFIG_REGISTRY[pipeline_name_or_path]
if pipeline_name_or_path in PIPE_NAME_TO_CONFIG:
pipeline_config_cls = PIPE_NAME_TO_CONFIG[pipeline_name_or_path]
# Try partial matches (for local paths that might include the weight ID)
for registered_id, config_class in WEIGHT_CONFIG_REGISTRY.items():
for registered_id, config_class in PIPE_NAME_TO_CONFIG.items():
if registered_id in pipeline_name_or_path:
return config_class
# If no match, try to use the fallback config
fallback_config = None
# Try to determine pipeline architecture for fallback
for pipeline_type, detector in PIPELINE_DETECTOR.items():
if detector(pipeline_name.lower()):
fallback_config = PIPELINE_FALLBACK_CONFIG.get(pipeline_type)
pipeline_config_cls = config_class
break
logger.warning("No match found for pipeline %s, using fallback config %s.",
pipeline_name_or_path, fallback_config)
return fallback_config
# If no match, try to use the fallback config
if pipeline_config_cls is None:
if os.path.exists(pipeline_name_or_path):
config = verify_model_config_and_directory(pipeline_name_or_path)
else:
config = maybe_download_model_index(pipeline_name_or_path)
logger.warning(
"Trying to use the config from the model_index.json. FastVideo may not correctly identify the optimal config for this model in this situation."
)
pipeline_name = config["_class_name"]
# Try to determine pipeline architecture for fallback
for pipeline_type, detector in PIPELINE_DETECTOR.items():
if detector(pipeline_name.lower()):
pipeline_config_cls = PIPELINE_FALLBACK_CONFIG.get(
pipeline_type)
break
if pipeline_config_cls is not None:
logger.warning(
"No match found for pipeline %s, using fallback config %s.",
pipeline_name_or_path, pipeline_config_cls)
if pipeline_config_cls is None:
raise ValueError(
f"No match found for pipeline {pipeline_name_or_path}, please check the pipeline name or path."
)
return pipeline_config_cls
-7
View File
@@ -39,7 +39,6 @@ class SamplingParam:
num_inference_steps: int = 50
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
VSA_sparsity: float = 0.0
# TeaCache parameters
enable_teacache: bool = False
@@ -185,12 +184,6 @@ class SamplingParam:
default=SamplingParam.image_path,
help="Path to input image for image-to-video generation",
)
parser.add_argument(
"--VSA-sparsity",
type=float,
default=SamplingParam.VSA_sparsity,
help="VSA attention sparsity",
)
return parser
+45
View File
@@ -0,0 +1,45 @@
from typing import Any, Dict
def update_config_from_args(config: Any,
args_dict: Dict[str, Any],
prefix: str = "",
pop_args: bool = False) -> None:
"""
Update configuration object from arguments dictionary.
Args:
config: The configuration object to update
args_dict: Dictionary containing arguments
prefix: Prefix for the configuration parameters in the args_dict.
If None, assumes direct attribute mapping without prefix.
"""
# Handle top-level attributes (no prefix)
args_not_to_remove = [
'model_path',
]
args_to_remove = []
if prefix.strip() == "":
for key, value in args_dict.items():
if hasattr(config, key) and value is not None:
if key == "text_encoder_precisions" and isinstance(value, list):
setattr(config, key, tuple(value))
else:
setattr(config, key, value)
if pop_args:
args_to_remove.append(key)
else:
# Handle nested attributes with prefix
prefix_with_dot = f"{prefix}."
for key, value in args_dict.items():
if key.startswith(prefix_with_dot) and value is not None:
attr_name = key[len(prefix_with_dot):]
if hasattr(config, attr_name):
setattr(config, attr_name, value)
if pop_args:
args_to_remove.append(key)
if pop_args:
for key in args_to_remove:
if key not in args_not_to_remove:
args_dict.pop(key)
+17 -20
View File
@@ -1,19 +1,17 @@
import os
# SPDX-License-Identifier: Apache-2.0
from torchvision import transforms
from torchvision.transforms import Lambda
from transformers import AutoTokenizer
from fastvideo.v1.dataset.t2v_datasets import T2V_dataset
from fastvideo.v1.dataset.parquet_dataset_map_style import (
build_parquet_map_style_dataloader)
from fastvideo.v1.dataset.preprocessing_datasets import (
VideoCaptionMergedDataset)
from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
from .parquet_dataset_map_style import build_parquet_map_style_dataloader
__all__ = ["build_parquet_map_style_dataloader"]
from fastvideo.v1.dataset.validation_dataset import ValidationDataset
def getdataset(args, start_idx=0) -> T2V_dataset:
def getdataset(args) -> VideoCaptionMergedDataset:
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
resize_topcrop = [
@@ -31,15 +29,14 @@ def getdataset(args, start_idx=0) -> T2V_dataset:
*resize_topcrop,
norm_fun,
])
tokenizer_path = os.path.join(args.model_path, "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
if args.dataset == "t2v":
return T2V_dataset(args,
transform=transform,
temporal_sample=temporal_sample,
tokenizer=tokenizer,
transform_topcrop=transform_topcrop,
start_idx=start_idx)
return VideoCaptionMergedDataset(data_merge_path=args.data_merge_path,
args=args,
transform=transform,
temporal_sample=temporal_sample,
transform_topcrop=transform_topcrop)
raise NotImplementedError(args.dataset)
__all__ = [
"build_parquet_map_style_dataloader", "ValidationDataset",
"VideoCaptionMergedDataset"
]
@@ -0,0 +1,185 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import os
import pathlib
import time
import torch.distributed as dist
import torch.distributed.checkpoint as dist_cp
from fastvideo.v1.dataset.parquet_dataset_iterable_style import (
build_parquet_iterable_style_dataloader)
from fastvideo.v1.distributed import get_world_rank
from fastvideo.v1.distributed.parallel_state import (
cleanup_dist_env_and_memory, get_local_torch_device,
maybe_init_distributed_environment_and_model_parallel)
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
def main() -> None:
parser = argparse.ArgumentParser(
description="Benchmark parquet iterable style dataset loading speed")
parser.add_argument(
"--path",
type=str,
help="Path to parquet dataset",
)
parser.add_argument("--batch_size",
type=int,
default=4,
help="Batch size for DataLoader")
parser.add_argument("--num_data_workers",
type=int,
help="Number of DataLoader workers")
parser.add_argument("--num_epoch",
type=int,
default=2,
help="Number of epoches to benchmark")
parser.add_argument("--verify_resume",
action="store_true",
help="Verify resume")
parser.add_argument(
"--num_batches_per_epoch",
type=int,
default=1000,
help="Number of batches to benchmark",
)
parser.add_argument('--checkpoint_path',
type=str,
default='dataloader_checkpoint',
help='Path to save/load checkpoint')
'''
example launch command:
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 2 --num_epoch 2 --num_batches_per_epoch 2 --verify_resume
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5 --verify_resume
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path /mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents/ --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 100
'''
args = parser.parse_args()
world_size = int(os.environ.get("WORLD_SIZE", 1))
maybe_init_distributed_environment_and_model_parallel(
tp_size=(world_size + 1) // 2, sp_size=(world_size + 1) // 2)
logger.info("Initialized distributed environment with world_size=%d",
world_size)
# Create DataLoader with proper settings
dataset, dataloader = build_parquet_iterable_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
logger.info("Initialized dataloader")
if args.verify_resume:
# First pass - record latent sums
first_pass_sums = []
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
latent_sum = latents.sum().item()
first_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f", i, latent_sum)
if i >= args.num_batches_per_epoch - 1:
break
# Save dataloader state using distributed checkpoint
checkpoint_dir = pathlib.Path(args.checkpoint_path)
logger.info("Rank %d: Saving dataloader state to %s", get_world_rank(),
checkpoint_dir)
states = {"dataloader": dataloader}
begin_time = time.monotonic()
dist_cp.save(states, checkpoint_id=checkpoint_dir.as_posix())
end_time = time.monotonic()
logger.info("Rank %d: Saved checkpoint in %.2f seconds",
get_world_rank(), end_time - begin_time)
# Make sure all processes wait for checkpoint to be saved
if world_size > 1:
dist.barrier()
# Recreate dataloader and load state
dataset, dataloader = build_parquet_iterable_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
load_states = {"dataloader": dataloader}
dist_cp.load(load_states, checkpoint_id=checkpoint_dir.as_posix())
logger.info("Rank %d: Loaded dataloader state from %s",
get_world_rank(), checkpoint_dir)
# Second pass - verify latent sums match
for i, (latents, embeddings, masks) in enumerate(dataloader):
latent_sum = latents.sum().item()
first_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f",
i + args.num_batches_per_epoch, latent_sum)
if i >= args.num_batches_per_epoch - 1:
break
dataset, dataloader = build_parquet_iterable_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
# Second pass - verify latent sums match
second_pass_sums = []
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
latent_sum = latents.sum().item()
second_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f (should match first pass: %f)",
i, latent_sum, first_pass_sums[i])
if i >= args.num_batches_per_epoch * 2 - 1:
break
# Verify all sums match
if all(
abs(a - b) < 1e-6
for a, b in zip(first_pass_sums, second_pass_sums)):
logger.info(
"All latent sums match between passes - resume verification successful!"
)
else:
raise ValueError(
"Latent sums do not match between passes - resume verification failed!"
)
start_time = time.time()
total_samples = 0
total_batches = 0
for _ in range(args.num_epoch):
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
if i >= args.num_batches_per_epoch:
break
# Move data to device
latents = latents.to(get_local_torch_device())
embeddings = embeddings.to(get_local_torch_device())
# Calculate actual batch size
batch_size = latents.size(0)
total_samples += batch_size
total_batches += 1
# Print progress only from rank 0
if get_world_rank() == 0 and (i + 1) % 10 == 0:
elapsed = time.time() - start_time
samples_per_sec = total_samples / elapsed
logger.info("Batch %d/%d, Speed: %.2f samples/sec", i + 1,
args.num_batches_per_epoch, samples_per_sec)
# Final statistics
if world_size > 1:
dist.barrier()
if get_world_rank() == 0:
elapsed = time.time() - start_time
samples_per_sec = total_samples / elapsed
logger.info("\nBenchmark Results:")
logger.info("Total time: %.2f seconds", elapsed)
logger.info("Total samples: %d", total_samples)
logger.info("Average speed: %.2f samples/sec", samples_per_sec)
logger.info("Time per batch: %.2f ms", elapsed / total_batches * 1000)
if __name__ == "__main__":
try:
main()
finally:
cleanup_dist_env_and_memory()
@@ -8,11 +8,12 @@ import torch
import torch.distributed as dist
import torch.distributed.checkpoint as dist_cp
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.v1.dataset.parquet_dataset_map_style import (
build_parquet_map_style_dataloader)
from fastvideo.v1.distributed import get_world_rank
from fastvideo.v1.distributed.parallel_state import (
cleanup_dist_env_and_memory, get_torch_device,
cleanup_dist_env_and_memory, get_local_torch_device,
maybe_init_distributed_environment_and_model_parallel)
from fastvideo.v1.logger import init_logger
@@ -54,9 +55,9 @@ def main() -> None:
help='Path to save/load checkpoint')
'''
example launch command:
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 2 --verify_resume
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 3 --verify_resume
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5 --verify_resume
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path /mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents/ --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 100
'''
args = parser.parse_args()
world_size = int(os.environ.get("WORLD_SIZE", 1))
@@ -66,14 +67,22 @@ def main() -> None:
world_size)
# Create DataLoader with proper settings
dataloader = build_parquet_map_style_dataloader(args.path, args.batch_size,
args.num_data_workers)
dataset, dataloader = build_parquet_map_style_dataloader(
args.path,
args.batch_size,
parquet_schema=pyarrow_schema_t2v,
num_data_workers=args.num_data_workers)
logger.info("Initialized dataloader with %d batches", len(dataloader))
if args.verify_resume:
for i, (latents, embeddings, masks,
data_indices) in enumerate(dataloader):
logger.info("Batch %d data_indices: %s", i, data_indices)
# First pass - record latent sums
first_pass_sums = []
for i, batch in enumerate(dataloader):
latents = batch['vae_latent']
embeddings = batch['text_embedding']
latent_sum = latents.sum().item()
first_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f", i, latent_sum)
if i >= args.num_batches_per_epoch - 1:
break
@@ -94,47 +103,70 @@ def main() -> None:
if world_size > 1:
dist.barrier()
dataloader = build_parquet_map_style_dataloader(args.path,
args.batch_size,
args.num_data_workers)
# Load dataloader state using distributed checkpoint
logger.info("Rank %d: Loading dataloader state from %s",
get_world_rank(), checkpoint_dir)
# Recreate dataloader and load state
dataset, dataloader = build_parquet_map_style_dataloader(
args.path,
args.batch_size,
parquet_schema=pyarrow_schema_t2v,
num_data_workers=args.num_data_workers)
load_states = {"dataloader": dataloader}
dist_cp.load(load_states, checkpoint_id=checkpoint_dir.as_posix())
logger.info("Rank %d: Loaded dataloader state from %s",
get_world_rank(), checkpoint_dir)
for i, (latents, embeddings, masks,
data_indices) in enumerate(dataloader):
logger.info("Batch %d data_indices: %s", i, data_indices)
for i, batch in enumerate(dataloader):
latents = batch['vae_latent']
embeddings = batch['text_embedding']
latent_sum = latents.sum().item()
first_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f",
i + args.num_batches_per_epoch, latent_sum)
if i >= args.num_batches_per_epoch - 1:
break
logger.info("Restart from the beginning")
dataset, dataloader = build_parquet_map_style_dataloader(
args.path,
args.batch_size,
parquet_schema=pyarrow_schema_t2v,
num_data_workers=args.num_data_workers)
dataloader = build_parquet_map_style_dataloader(args.path,
args.batch_size,
args.num_data_workers)
for i, (latents, embeddings, masks,
data_indices) in enumerate(dataloader):
logger.info("Batch %d data_indices: %s", i, data_indices)
# Second pass - verify latent sums match
second_pass_sums = []
for i, batch in enumerate(dataloader):
latents = batch['vae_latent']
embeddings = batch['text_embedding']
latent_sum = latents.sum().item()
second_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f (should match first pass: %f)",
i, latent_sum, first_pass_sums[i])
if i >= args.num_batches_per_epoch * 2 - 1:
break
# Verify all sums match
if all(
abs(a - b) < 1e-6
for a, b in zip(first_pass_sums, second_pass_sums)):
logger.info(
"All latent sums match between passes - resume verification successful!"
)
else:
raise ValueError(
"Latent sums do not match between passes - resume verification failed!"
)
start_time = time.time()
total_samples = 0
total_batches = 0
for _ in range(args.num_epoch):
for i, (latents, embeddings, masks,
data_indices) in enumerate(dataloader):
for i, batch in enumerate(dataloader):
latents = batch['vae_latent']
embeddings = batch['text_embedding']
if i >= args.num_batches_per_epoch:
break
# Move data to device
latents = latents.to(get_torch_device())
embeddings = embeddings.to(get_torch_device())
latents = latents.to(get_local_torch_device())
embeddings = embeddings.to(get_local_torch_device())
# Calculate actual batch size
batch_size = latents.size(0)
+60 -11
View File
@@ -26,15 +26,47 @@ pyarrow_schema_i2v = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
#I2V
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_dtype", pa.string()),
pa.field("first_frame_latent_bytes", pa.binary()),
pa.field("first_frame_latent_shape", pa.list_(pa.int64())),
pa.field("first_frame_latent_dtype", pa.string()),
# I2V Validation
pa.field("pil_image_bytes", pa.binary()),
pa.field("pil_image_shape", pa.list_(pa.int64())),
pa.field("pil_image_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
pyarrow_schema_i2v_validation = pa.schema([
pa.field("id", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
# e.g., [SeqLen, Dim]
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
#I2V
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_dtype", pa.string()),
# I2V Validation
pa.field("pil_image_bytes", pa.binary()),
pa.field("pil_image_shape", pa.list_(pa.int64())),
pa.field("pil_image_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
@@ -64,11 +96,6 @@ pyarrow_schema_t2v = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_attention_mask_bytes", pa.binary()),
# e.g., [SeqLen]
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
# e.g., 'bool' or 'int8'
pa.field("text_attention_mask_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
@@ -80,4 +107,26 @@ pyarrow_schema_t2v = pa.schema([
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
])
pyarrow_schema_t2v_validation = pa.schema([
pa.field("id", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
# e.g., [SeqLen, Dim]
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
-137
View File
@@ -1,137 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import json
import os
import time
from multiprocessing import Pool, cpu_count
from pathlib import Path
import torchvision
from tqdm import tqdm
def get_video_info(video_path):
"""Get video information using torchvision."""
# Read video tensor (T, C, H, W)
video_tensor, _, info = torchvision.io.read_video(str(video_path),
output_format="TCHW",
pts_unit="sec")
num_frames = video_tensor.shape[0]
height = video_tensor.shape[2]
width = video_tensor.shape[3]
fps = info.get("video_fps", 0)
duration = num_frames / fps if fps > 0 else 0
# Extract name
_, _, videos_dir, video_name = str(video_path).split("/")
return {
"path": str(video_name),
"resolution": {
"width": width,
"height": height
},
"size": os.path.getsize(video_path),
"fps": fps,
"duration": duration,
"num_frames": num_frames
}
def prepare_dataset_json(folder_path,
output_name="videos2caption.json",
num_workers=None) -> None:
"""Prepare dataset information from a folder containing videos and prompt.txt."""
folder_path = Path(folder_path)
# Read prompt file
prompt_file = folder_path / "prompt.txt"
if not prompt_file.exists():
raise FileNotFoundError(f"prompt.txt not found in {folder_path}")
with open(prompt_file) as f:
prompts = [line.strip() for line in f.readlines() if line.strip()]
# Read videos file
videos_file = folder_path / "videos.txt"
if not videos_file.exists():
raise FileNotFoundError(f"videos.txt not found in {folder_path}")
with open(videos_file) as f:
video_paths = [line.strip() for line in f.readlines() if line.strip()]
if len(prompts) != len(video_paths):
raise ValueError(
f"Number of prompts ({len(prompts)}) does not match number of videos ({len(video_paths)})"
)
# Prepare arguments for multiprocessing
process_args = [folder_path / video_path for video_path in video_paths]
# Determine number of workers
if num_workers is None:
num_workers = max(1, cpu_count() - 1) # Leave one CPU free
# Process videos in parallel
start_time = time.time()
with Pool(num_workers) as pool:
results = list(
tqdm(pool.imap(get_video_info, process_args),
total=len(process_args),
desc="Processing videos",
unit="video"))
# Combine results with prompts
dataset_info = []
for result, prompt in zip(results, prompts):
result["cap"] = [prompt]
dataset_info.append(result)
# Calculate total processing time
total_time = time.time() - start_time
total_videos = len(dataset_info)
avg_time_per_video = total_time / total_videos if total_videos > 0 else 0
print("\nProcessing completed:")
print(f"Total videos processed: {total_videos}")
print(f"Total time: {total_time:.2f} seconds")
print(f"Average time per video: {avg_time_per_video:.2f} seconds")
# Save to JSON file
output_file = folder_path / output_name
with open(output_file, 'w') as f:
json.dump(dataset_info, f, indent=2)
# Create merge.txt
merge_file = folder_path / "merge.txt"
with open(merge_file, 'w') as f:
f.write(f"{folder_path}/videos,{output_file}\n")
print(f"Dataset information saved to {output_file}")
print(f"Merge file created at {merge_file}")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description='Prepare video dataset information in JSON format')
parser.add_argument(
'--folder',
type=str,
required=True,
help='Path to the folder containing videos and prompt.txt')
parser.add_argument(
'--output',
type=str,
default='videos2caption.json',
help='Name of the output JSON file (default: videos2caption.json)')
parser.add_argument('--workers',
type=int,
default=32,
help='Number of worker processes (default: 16)')
return parser.parse_args()
if __name__ == "__main__":
args = parse_args()
prepare_dataset_json(args.folder, args.output, args.workers)
@@ -0,0 +1,278 @@
import os
import pickle
import random
from typing import Dict, List, Tuple
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
import tqdm
from torch.utils.data import IterableDataset, get_worker_info
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.dataset.utils import collate_latents_embs_masks
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
get_world_size)
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class BatchIterator:
# TODO: Implement state_dict and load_state_dict to support resume.
def __init__(self, files, batch_size, text_padding_length, keys,
worker_num_samples, read_batch_size):
self.files = files
self.batch_size = batch_size
self.text_padding_length = text_padding_length
self.keys = keys
self.worker_num_samples = worker_num_samples
self.processed_samples = 0
self.buffer = []
self.read_batch_size = read_batch_size
def __iter__(self):
for file in self.files:
if self.processed_samples >= self.worker_num_samples:
return
reader = pq.ParquetFile(file)
for batch in reader.iter_batches(batch_size=self.read_batch_size):
if self.processed_samples >= self.worker_num_samples:
return
self.buffer.extend(batch.to_pylist())
while len(self.buffer) >= self.batch_size:
if self.processed_samples >= self.worker_num_samples:
return
batch_to_process = self.buffer[:self.batch_size]
self.buffer = self.buffer[self.batch_size:]
all_latents, all_embs, all_masks, caption_text = collate_latents_embs_masks(
batch_to_process, self.text_padding_length, self.keys)
self.processed_samples += self.batch_size
yield all_latents, all_embs, all_masks, caption_text
class LatentsParquetIterStyleDataset(IterableDataset):
"""Efficient loader for video-text data from a directory of Parquet files."""
# Modify this in the future if we want to add more keys, for example, in image to video.
keys = [("vae_latent", "latent"), ("text_embedding")]
def __init__(self,
path: str,
batch_size: int = 1024,
cfg_rate: float = 0.1,
num_workers: int = 1,
drop_last: bool = True,
text_padding_length: int = 512,
seed: int = 42,
read_batch_size: int = 32,
parquet_schema: pa.Schema = None):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
self.parquet_schema = parquet_schema
self.cfg_rate = cfg_rate
self.text_padding_length = text_padding_length
self.seed = seed
self.read_batch_size = read_batch_size
# Get distributed training info
self.global_rank = get_world_rank()
self.world_size = get_world_size()
self.sp_world_size = get_sp_world_size()
self.num_sp_groups = self.world_size // self.sp_world_size
num_workers = 1 if num_workers == 0 else num_workers
# Get sharding info
shard_parquet_files, shard_total_samples, shard_parquet_lengths = shard_parquet_files_across_sp_groups_and_workers(
self.path, self.num_sp_groups, num_workers, seed)
if drop_last:
self.worker_num_samples = min(
shard_total_samples) // batch_size * batch_size
# Assign files to current rank's SP group
ith_sp_group = self.global_rank // self.sp_world_size
self.sp_group_parquet_files = shard_parquet_files[ith_sp_group::self
.num_sp_groups]
self.sp_group_parquet_lengths = shard_parquet_lengths[
ith_sp_group::self.num_sp_groups]
self.sp_group_num_samples = shard_total_samples[ith_sp_group::self.
num_sp_groups]
logger.info(
"In total %d parquet files, %d samples, after sharding we retain %d samples due to drop_last",
sum([len(shard) for shard in shard_parquet_files]),
sum(shard_total_samples),
self.worker_num_samples * self.num_sp_groups * num_workers)
else:
raise ValueError("drop_last must be True")
logger.info("Each dataloader worker will load %d samples",
self.worker_num_samples)
def __iter__(self):
worker_info = get_worker_info()
worker_id = worker_info.id if worker_info is not None else 1
worker_files = self.sp_group_parquet_files[worker_id]
batch_iterator = BatchIterator(
files=worker_files,
batch_size=self.batch_size,
text_padding_length=self.text_padding_length,
keys=self.keys,
worker_num_samples=self.worker_num_samples,
read_batch_size=self.read_batch_size) # type: ignore
yield from batch_iterator
if batch_iterator.processed_samples != self.worker_num_samples:
raise ValueError(
"Rank %d, Worker %d: Not enough samples to process, this should not happen",
self.global_rank, worker_id)
def shard_parquet_files_across_sp_groups_and_workers(
path: str,
num_sp_groups: int,
num_workers: int,
seed: int = 42,
) -> Tuple[List[List[str]], List[int], List[Dict[str, int]]]:
"""
Shard parquet files across SP groups and workers in a balanced way.
Args:
path: Directory containing parquet files
num_sp_groups: Number of SP groups to shard across
num_workers: Number of workers per SP group
seed: Random seed for shuffling
Returns:
Tuple containing:
- List of lists of parquet files for each shard
- List of total samples per shard
- List of dictionaries mapping file paths to their lengths
"""
# Check if sharding plan already exists
sharding_info_dir = os.path.join(
path, f"sharding_info_{num_sp_groups}_sp_groups_{num_workers}_workers")
if os.path.exists(sharding_info_dir):
logger.info("Sharding plan already exists")
logger.info("Loading sharding plan from %s", sharding_info_dir)
try:
with open(
os.path.join(sharding_info_dir, "shard_parquet_files.pkl"),
"rb") as f:
shard_parquet_files = pickle.load(f)
with open(
os.path.join(sharding_info_dir, "shard_total_samples.pkl"),
"rb") as f:
shard_total_samples = pickle.load(f)
with open(
os.path.join(sharding_info_dir,
"shard_parquet_lengths.pkl"), "rb") as f:
shard_parquet_lengths = pickle.load(f)
return shard_parquet_files, shard_total_samples, shard_parquet_lengths
except Exception as e:
logger.error("Error loading sharding plan: %s", str(e))
logger.info("Falling back to creating new sharding plan")
if get_world_rank() == 0:
logger.info("Scanning for parquet files in %s", path)
# Find all parquet files
parquet_files = []
for root, _, files in os.walk(path):
for file in files:
if file.endswith('.parquet'):
parquet_files.append(os.path.join(root, file))
if not parquet_files:
raise ValueError("No parquet files found in %s", path)
# Calculate file lengths efficiently using a single pass
logger.info("Calculating file lengths...")
lengths = []
for file in tqdm.tqdm(parquet_files, desc="Reading parquet files"):
lengths.append(pq.ParquetFile(file).metadata.num_rows)
total_samples = sum(lengths)
logger.info("Found %d files with %d total samples", len(parquet_files),
total_samples)
# Sort files by length for better balancing
sorted_indices = np.argsort(lengths)
sorted_files = [parquet_files[i] for i in sorted_indices]
sorted_lengths = [lengths[i] for i in sorted_indices]
# Create shards
num_shards = num_sp_groups * num_workers
shard_parquet_files = [[] for _ in range(num_shards)]
shard_total_samples = [0] * num_shards
shard_parquet_lengths = [{} for _ in range(num_shards)]
# Distribute files to shards using a greedy approach
logger.info("Distributing files to shards...")
for file, length in zip(reversed(sorted_files),
reversed(sorted_lengths)):
# Find shard with minimum current length
target_shard = np.argmin(shard_total_samples)
shard_parquet_files[target_shard].append(file)
shard_total_samples[target_shard] += length
shard_parquet_lengths[target_shard][file] = length
#randomize each shard
for shard in shard_parquet_files:
random.seed(seed)
random.shuffle(shard)
save_dir = os.path.join(
path,
f"sharding_info_{num_sp_groups}_sp_groups_{num_workers}_workers")
os.makedirs(save_dir, exist_ok=True)
with open(os.path.join(save_dir, "shard_parquet_files.pkl"), "wb") as f:
pickle.dump(shard_parquet_files, f)
with open(os.path.join(save_dir, "shard_total_samples.pkl"), "wb") as f:
pickle.dump(shard_total_samples, f)
with open(os.path.join(save_dir, "shard_parquet_lengths.pkl"),
"wb") as f:
pickle.dump(shard_parquet_lengths, f)
logger.info("Saved sharding info to %s", save_dir)
# wait for all ranks to finish
torch.distributed.barrier()
# recursive call
return shard_parquet_files_across_sp_groups_and_workers(
path, num_sp_groups, num_workers, seed)
def build_parquet_iterable_style_dataloader(
path: str,
batch_size: int,
num_data_workers: int,
cfg_rate: float = 0.0,
drop_last: bool = True,
text_padding_length: int = 512,
seed: int = 42,
read_batch_size: int = 32
) -> Tuple[LatentsParquetIterStyleDataset, StatefulDataLoader]:
"""Build a dataloader for the LatentsParquetIterStyleDataset."""
dataset = LatentsParquetIterStyleDataset(
path=path,
batch_size=batch_size,
cfg_rate=cfg_rate,
num_workers=num_data_workers,
drop_last=drop_last,
text_padding_length=text_padding_length,
seed=seed,
read_batch_size=read_batch_size)
loader = StatefulDataLoader(
dataset,
batch_size=1,
num_workers=num_data_workers,
pin_memory=True,
)
return dataset, loader
+92 -106
View File
@@ -1,15 +1,18 @@
# SPDX-License-Identifier: Apache-2.0
import os
import pickle
from typing import Any, Dict, List, Tuple
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
# Torch in general
import torch
import tqdm
# Dataset
from torch.utils.data import Dataset, Sampler
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.dataset.utils import collate_rows_from_parquet_schema
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
get_world_size)
from fastvideo.v1.logger import init_logger
@@ -30,6 +33,7 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
sp_world_size: int,
global_rank: int,
drop_last: bool = True,
drop_first_row: bool = False,
seed: int = 0,
):
self.batch_size = batch_size
@@ -45,6 +49,11 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
# Create a random permutation of all indices
global_indices = torch.randperm(self.dataset_size, generator=rng)
if drop_first_row:
# drop 0 in global_indices
global_indices = global_indices[global_indices != 0]
self.dataset_size = self.dataset_size - 1
if self.drop_last:
# For drop_last=True, we:
# 1. Ensure total samples is divisible by (batch_size * num_sp_groups)
@@ -56,19 +65,22 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
self.num_sp_groups *
self.batch_size]
else:
# add more indices to make it divisible by (batch_size * num_sp_groups)
padding_size = self.num_sp_groups * self.batch_size - (
self.dataset_size % (self.num_sp_groups * self.batch_size))
global_indices = torch.cat(
[global_indices, global_indices[:padding_size]])
if self.dataset_size % (self.num_sp_groups * self.batch_size) != 0:
# add more indices to make it divisible by (batch_size * num_sp_groups)
padding_size = self.num_sp_groups * self.batch_size - (
self.dataset_size % (self.num_sp_groups * self.batch_size))
logger.info("Padding the dataset from %d to %d",
self.dataset_size, self.dataset_size + padding_size)
global_indices = torch.cat(
[global_indices, global_indices[:padding_size]])
# shard the indices to each sp group
ith_sp_group = self.global_rank // self.sp_world_size
sp_group_local_indices = global_indices[ith_sp_group::self.
num_sp_groups]
self.sp_group_local_indices = sp_group_local_indices
logger.info("sp_group_local_indices: %d", len(sp_group_local_indices))
logger.info("Dataset size for each sp group: %d",
len(sp_group_local_indices))
def __iter__(self):
indices = self.sp_group_local_indices
@@ -81,20 +93,49 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
def get_parquet_files_and_length(path: str):
lengths = []
file_names = []
for root, _, files in os.walk(path):
for file in sorted(files):
if file.endswith('.parquet'):
file_path = os.path.join(root, file)
num_rows = pq.ParquetFile(file_path).metadata.num_rows
lengths.append(num_rows)
file_names.append(file_path)
# sort according to file name to ensure all rank has the same order (in case os.walk is not sorted)
file_names_sorted, lengths_sorted = zip(
*sorted(zip(file_names, lengths), key=lambda x: x[0]))
assert len(file_names_sorted) != 0, "No parquet files found in the dataset"
return file_names_sorted, lengths_sorted
# Check if cached info exists
cache_dir = os.path.join(path, "map_style_cache")
cache_file = os.path.join(cache_dir, "file_info.pkl")
if os.path.exists(cache_file):
logger.info("Loading cached file info from %s", cache_file)
try:
with open(cache_file, "rb") as f:
file_names_sorted, lengths_sorted = pickle.load(f)
return file_names_sorted, lengths_sorted
except Exception as e:
logger.error("Error loading cached file info: %s", str(e))
logger.info("Falling back to scanning files")
# If no cache exists or loading failed, scan files
if get_world_rank() == 0:
lengths = []
file_names = []
for root, _, files in os.walk(path):
for file in sorted(files):
if file.endswith('.parquet'):
file_path = os.path.join(root, file)
file_names.append(file_path)
for file_path in tqdm.tqdm(file_names,
desc="Reading parquet files to get lengths"):
num_rows = pq.ParquetFile(file_path).metadata.num_rows
lengths.append(num_rows)
# sort according to file name to ensure all rank has the same order (in case os.walk is not sorted)
file_names_sorted, lengths_sorted = zip(
*sorted(zip(file_names, lengths), key=lambda x: x[0]))
assert len(
file_names_sorted) != 0, "No parquet files found in the dataset"
os.makedirs(cache_dir, exist_ok=True)
with open(cache_file, "wb") as f:
pickle.dump((file_names_sorted, lengths_sorted), f)
logger.info("Saved file info to %s", cache_file)
# Wait for rank 0 to finish saving
if get_world_size() > 1:
torch.distributed.barrier()
return get_parquet_files_and_length(path)
def read_row_from_parquet_file(parquet_files: List[str], global_row_idx: int,
@@ -144,39 +185,27 @@ class LatentsParquetMapStyleDataset(Dataset):
Note:
Using parquet for map style dataset is not efficient, we mainly keep it for backward compatibility and debugging.
"""
# Modify this in the future if we want to add more keys, for example, in image to video.
keys = ["vae_latent", "text_embedding"]
def __init__(
self,
path: str,
batch_size: int,
parquet_schema: pa.Schema,
cfg_rate: float = 0.0,
seed: int = 42,
drop_last: bool = True,
drop_first_row: bool = False,
text_padding_length: int = 512,
):
super().__init__()
self.path = path
self.cfg_rate = cfg_rate
if cfg_rate > 0.0:
raise ValueError(
"cfg_rate > 0.0 is not supported for now because it will trigger bug when num_data_workers > 0"
)
self.parquet_schema = parquet_schema
logger.info("Initializing LatentsParquetMapStyleDataset with path: %s",
path)
self.parquet_files, self.lengths = get_parquet_files_and_length(path)
self.batch = batch_size
self.text_padding_length = text_padding_length
self._cols = [
"vae_latent_bytes",
"vae_latent_shape",
"text_embedding_bytes",
"text_embedding_shape",
"text_embedding_dtype",
"height",
"width",
]
self.sampler = DP_SP_BatchSampler(
batch_size=batch_size,
dataset_size=sum(self.lengths),
@@ -184,27 +213,14 @@ class LatentsParquetMapStyleDataset(Dataset):
sp_world_size=get_sp_world_size(),
global_rank=get_world_rank(),
drop_last=drop_last,
drop_first_row=drop_first_row,
seed=seed,
)
logger.info("Dataset initialized with %d parquet files and %d rows",
len(self.parquet_files), sum(self.lengths))
def _get_torch_tensors_from_row_dict(
self, row_dict: Dict[str, Any]) -> Dict[str, torch.Tensor]:
"""
Get the latents and prompts from a row dictionary.
"""
return_dict = {}
for key in self.keys:
shape = row_dict[f"{key}_shape"]
bytes = row_dict[f"{key}_bytes"]
# TODO (peiyuan): read precision
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
data = torch.from_numpy(data)
return_dict[key] = data
return return_dict
def get_validation_negative_prompt(self) -> tuple[Any, Any, Any, Any]:
def get_validation_negative_prompt(
self) -> tuple[torch.Tensor, torch.Tensor, str]:
"""
Get the negative prompt for validation.
This method ensures the negative prompt is loaded and cached properly.
@@ -218,39 +234,23 @@ class LatentsParquetMapStyleDataset(Dataset):
row_dict = read_row_from_parquet_file([file_path], row_idx,
[self.lengths[0]])
# Get tensors using the existing helper method
data = self._get_torch_tensors_from_row_dict(row_dict)
emb = data["text_embedding"]
batch = collate_rows_from_parquet_schema([row_dict],
self.parquet_schema,
self.text_padding_length,
cfg_rate=0.0)
negative_prompt = batch['info_list'][0]['prompt']
negative_prompt_embedding = batch['text_embedding']
negative_prompt_attention_mask = batch['text_attention_mask']
if len(negative_prompt_embedding.shape) == 2:
negative_prompt_embedding = negative_prompt_embedding.unsqueeze(0)
if len(negative_prompt_attention_mask.shape) == 1:
negative_prompt_attention_mask = negative_prompt_attention_mask.unsqueeze(
0).unsqueeze(0)
# Pad the embedding and get mask
padded_emb, mask = self._pad(emb, self.text_padding_length)
# Pin memory for faster transfer to GPU
padded_emb = padded_emb
mask = mask
return None, padded_emb, mask, None
def _pad(self, t: torch.Tensor, padding_length: int) -> torch.Tensor:
"""
Pad or crop an embedding [L, D] to exactly padding_length tokens.
Return:
- [L, D] tensor in pinned CPU memory
- [L] attention mask in pinned CPU memory
"""
L, D = t.shape
if padding_length > L: # pad
pad = torch.zeros(padding_length - L,
D,
dtype=t.dtype,
device=t.device)
return torch.cat([t, pad], 0), torch.cat(
[torch.ones(L), torch.zeros(padding_length - L)], 0)
else: # crop
return t[:padding_length], torch.ones(padding_length)
return negative_prompt_embedding, negative_prompt_attention_mask, negative_prompt
# PyTorch calls this ONLY because the batch_sampler yields a list
def __getitems__(self, indices: List[int]):
def __getitems__(self, indices: List[int]) -> Dict[str, Any]:
"""
Batch fetch using read_row_from_parquet_file for each index.
"""
@@ -259,29 +259,11 @@ class LatentsParquetMapStyleDataset(Dataset):
for idx in indices
]
# Initialize tensors to hold padded embeddings and masks
all_latents = []
all_embs = []
all_masks = []
# Process each row individually
for i, row in enumerate(rows):
# Get tensors from row
data = self._get_torch_tensors_from_row_dict(row)
latents, emb = data["vae_latent"], data["text_embedding"]
padded_emb, mask = self._pad(emb, self.text_padding_length)
# Store in batch tensors
all_latents.append(latents)
all_embs.append(padded_emb)
all_masks.append(mask)
# Pin memory for faster transfer to GPU
all_latents = torch.stack(all_latents)
all_embs = torch.stack(all_embs)
all_masks = torch.stack(all_masks)
return all_latents, all_embs, all_masks, indices
batch = collate_rows_from_parquet_schema(rows,
self.parquet_schema,
self.text_padding_length,
cfg_rate=self.cfg_rate)
return batch
def __len__(self):
return sum(self.lengths)
@@ -298,8 +280,10 @@ def build_parquet_map_style_dataloader(
path,
batch_size,
num_data_workers,
parquet_schema,
cfg_rate=0.0,
drop_last=True,
drop_first_row=False,
text_padding_length=512,
seed=42) -> Tuple[LatentsParquetMapStyleDataset, StatefulDataLoader]:
dataset = LatentsParquetMapStyleDataset(
@@ -307,7 +291,9 @@ def build_parquet_map_style_dataloader(
batch_size,
cfg_rate=cfg_rate,
drop_last=drop_last,
drop_first_row=drop_first_row,
text_padding_length=text_padding_length,
parquet_schema=parquet_schema,
seed=seed)
loader = StatefulDataLoader(
@@ -0,0 +1,615 @@
# SPDX-License-Identifier: Apache-2.0
import json
import math
import os
import random
from abc import ABC, abstractmethod
from collections import Counter
from dataclasses import dataclass
from os.path import join as opj
from typing import Any, Dict, List, Optional, Union
import numpy as np
import torch
import torchvision
from einops import rearrange
from PIL import Image
from transformers import AutoTokenizer
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
@dataclass
class PreprocessBatch:
"""
Batch information for dataset processing stages.
This class holds all the information about a video-caption or image-caption pair
as it moves through the processing pipeline. Fields are populated by different stages.
"""
# Raw metadata
path: str
cap: Union[str, List[str]]
resolution: Optional[Dict] = None
fps: Optional[float] = None
duration: Optional[float] = None
# Processed metadata
num_frames: Optional[int] = None
sample_frame_index: Optional[List[int]] = None
sample_num_frames: Optional[int] = None
# Processed data
pixel_values: Optional[torch.Tensor] = None
text: Optional[str] = None
input_ids: Optional[torch.Tensor] = None
cond_mask: Optional[torch.Tensor] = None
@property
def is_video(self) -> bool:
"""Check if this is a video item."""
return self.path.endswith(".mp4")
@property
def is_image(self) -> bool:
"""Check if this is an image item."""
return self.path.endswith(".jpg")
class DatasetStage(ABC):
"""
Abstract base class for dataset processing stages.
Similar to PipelineStage but designed for dataset preprocessing operations.
"""
@abstractmethod
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Process the dataset batch.
Args:
batch: Dataset batch to process
**kwargs: Additional processing parameters
Returns:
Processed batch
"""
raise NotImplementedError
class DatasetFilterStage(ABC):
"""
Abstract base class for dataset filtering stages.
These stages can filter out items during metadata processing.
"""
@abstractmethod
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Check if batch should be kept.
Args:
batch: Dataset batch to check
**kwargs: Additional parameters
Returns:
True if batch should be kept, False otherwise
"""
raise NotImplementedError
@abstractmethod
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Process the dataset batch (for non-filtering operations).
Args:
batch: Dataset batch to process
**kwargs: Additional processing parameters
Returns:
Processed batch
"""
raise NotImplementedError
class DataValidationStage(DatasetFilterStage):
"""Stage for validating data items."""
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Validate data item.
Args:
batch: Dataset batch to validate
Returns:
True if valid, False if invalid
"""
# Check for caption
if batch.cap is None:
return False
if batch.is_video:
# Validate video-specific fields
if batch.duration is None or batch.fps is None:
return False
elif not batch.is_image:
return False
return True
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""Process does nothing for validation - filtering is handled by should_keep."""
return batch
class ResolutionFilterStage(DatasetFilterStage):
"""Stage for filtering data items based on resolution constraints."""
def __init__(self,
max_h_div_w_ratio: float = 17 / 16,
min_h_div_w_ratio: float = 8 / 16,
max_height: int = 1024,
max_width: int = 1024):
self.max_h_div_w_ratio = max_h_div_w_ratio
self.min_h_div_w_ratio = min_h_div_w_ratio
self.max_height = max_height
self.max_width = max_width
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Check if data item passes resolution filtering.
Args:
batch: Dataset batch with resolution information
Returns:
True if passes filter, False otherwise
"""
# Only apply to videos
if not batch.is_video:
return True
if batch.resolution is None:
return False
height = batch.resolution.get("height", None)
width = batch.resolution.get("width", None)
if height is None or width is None:
return False
# Check aspect ratio
aspect = self.max_height / self.max_width
hw_aspect_thr = 1.5
return self.filter_resolution(
height,
width,
max_h_div_w_ratio=hw_aspect_thr * aspect,
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
)
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""Process does nothing for resolution filtering - filtering is handled by should_keep."""
return batch
def filter_resolution(self, h: int, w: int, max_h_div_w_ratio: float,
min_h_div_w_ratio: float) -> bool:
"""Filter based on height/width ratio."""
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
class FrameSamplingStage(DatasetFilterStage):
"""Stage for temporal frame sampling and indexing."""
def __init__(self,
num_frames: int,
train_fps: int,
speed_factor: int = 1,
video_length_tolerance_range: float = 5.0,
drop_short_ratio: float = 0.0):
self.num_frames = num_frames
self.train_fps = train_fps
self.speed_factor = speed_factor
self.video_length_tolerance_range = video_length_tolerance_range
self.drop_short_ratio = drop_short_ratio
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
"""
Check if video should be kept based on length constraints.
Args:
batch: Dataset batch
Returns:
True if should be kept, False otherwise
"""
if batch.is_image:
return True
if batch.duration is None or batch.fps is None:
return False
num_frames = math.ceil(batch.fps * batch.duration)
# Check if video is too long
if (num_frames / batch.fps > self.video_length_tolerance_range *
(self.num_frames / self.train_fps * self.speed_factor)):
return False
# Resample frame indices to check length
frame_interval = batch.fps / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, num_frames,
frame_interval).astype(int)
# Filter short videos
return not (len(frame_indices) < self.num_frames
and random.random() < self.drop_short_ratio)
def process(self,
batch: PreprocessBatch,
temporal_sample_fn=None,
**kwargs) -> PreprocessBatch:
"""
Process frame sampling for video data items.
Args:
batch: Dataset batch
temporal_sample_fn: Function for temporal sampling
Returns:
Updated batch with frame sampling info
"""
if batch.is_image:
# For images, just add sample info
batch.sample_frame_index = [0]
batch.sample_num_frames = 1
return batch
assert batch.duration is not None and batch.fps is not None
batch.num_frames = math.ceil(batch.fps * batch.duration)
# Resample frame indices
frame_interval = batch.fps / self.train_fps
start_frame_idx = 0
frame_indices = np.arange(start_frame_idx, batch.num_frames,
frame_interval).astype(int)
# Temporal crop if too long
if len(frame_indices
) > self.num_frames and temporal_sample_fn is not None:
begin_index, end_index = temporal_sample_fn(len(frame_indices))
frame_indices = frame_indices[begin_index:end_index]
batch.sample_frame_index = frame_indices.tolist()
batch.sample_num_frames = len(frame_indices)
return batch
class VideoTransformStage(DatasetStage):
"""Stage for video data transformation."""
def __init__(self, transform) -> None:
self.transform = transform
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Transform video data.
Args:
batch: Dataset batch with video information
Returns:
Batch with transformed video tensor
"""
if not batch.is_video:
return batch
assert os.path.exists(batch.path), f"file {batch.path} do not exist!"
assert batch.sample_frame_index is not None, "Frame indices must be set before transformation"
torchvision_video, _, metadata = torchvision.io.read_video(
batch.path, output_format="TCHW")
video = torchvision_video[batch.sample_frame_index]
if self.transform is not None:
video = self.transform(video)
video = rearrange(video, "t c h w -> c t h w")
video = video.to(torch.uint8)
h, w = video.shape[-2:]
assert (
h / w <= 17 / 16 and h / w >= 8 / 16
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({batch.path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
video = video.float() / 127.5 - 1.0
batch.pixel_values = video
return batch
class ImageTransformStage(DatasetStage):
"""Stage for image data transformation."""
def __init__(self, transform, transform_topcrop) -> None:
self.transform = transform
self.transform_topcrop = transform_topcrop
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Transform image data.
Args:
batch: Dataset batch with image information
Returns:
Batch with transformed image tensor
"""
if not batch.is_image:
return batch
image = Image.open(batch.path).convert("RGB")
image = torch.from_numpy(np.array(image))
image = rearrange(image, "h w c -> c h w").unsqueeze(0)
if self.transform_topcrop is not None:
image = self.transform_topcrop(image)
elif self.transform is not None:
image = self.transform(image)
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
image = image.float() / 127.5 - 1.0
batch.pixel_values = image
return batch
class TextEncodingStage(DatasetStage):
"""Stage for text tokenization and encoding."""
def __init__(self, tokenizer, text_max_length: int, cfg_rate: float = 0.0):
self.tokenizer = tokenizer
self.text_max_length = text_max_length
self.cfg_rate = cfg_rate
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
"""
Process text data.
Args:
batch: Dataset batch with caption information
Returns:
Batch with encoded text information
"""
text = batch.cap
if not isinstance(text, list):
text = [text]
text = [random.choice(text)]
text = text[0] if random.random() > self.cfg_rate else ""
text_tokens_and_mask = self.tokenizer(
text,
max_length=self.text_max_length,
padding="max_length",
truncation=True,
return_attention_mask=True,
add_special_tokens=True,
return_tensors="pt",
)
batch.text = text
batch.input_ids = text_tokens_and_mask["input_ids"]
batch.cond_mask = text_tokens_and_mask["attention_mask"]
return batch
class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
torch.distributed.checkpoint.stateful.Stateful):
"""
Merged dataset for video and caption data with stage-based processing.
Assumes that data_merge_path is a txt file with the following format:
<folder_path>,<json_file_path>
The folder should contain videos.
The json file should be a list of dictionaries with the following format:
[
{
"path": "1gGQy4nxyUo-Scene-016.mp4",
"resolution": {
"width": 1920,
"height": 1080
},
"size": 2439112,
"fps": 25.0,
"duration": 6.88,
"num_frames": 172,
"cap": [
"A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open."
]
},
...
]
This dataset processes video and image data through a series of stages:
- Data validation
- Resolution filtering
- Frame sampling
- Transformation
- Text encoding
"""
def __init__(self,
data_merge_path: str,
args,
transform,
temporal_sample,
transform_topcrop,
start_idx: int = 0):
self.data_merge_path = data_merge_path
self.start_idx = start_idx
self.args = args
self.temporal_sample = temporal_sample
# Initialize tokenizer
tokenizer_path = os.path.join(args.model_path, "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
# Initialize processing stages
self._init_stages(args, transform, transform_topcrop, tokenizer)
# Process metadata
self.processed_batches = self._process_metadata()
def _init_stages(self, args, transform, transform_topcrop,
tokenizer) -> None:
"""Initialize all processing stages."""
self.validation_stage = DataValidationStage()
self.resolution_filter_stage = ResolutionFilterStage(
max_height=args.max_height, max_width=args.max_width)
self.frame_sampling_stage = FrameSamplingStage(
num_frames=args.num_frames,
train_fps=args.train_fps,
speed_factor=args.speed_factor,
video_length_tolerance_range=args.video_length_tolerance_range,
drop_short_ratio=args.drop_short_ratio)
self.video_transform_stage = VideoTransformStage(transform)
self.image_transform_stage = ImageTransformStage(
transform, transform_topcrop)
self.text_encoding_stage = TextEncodingStage(
tokenizer=tokenizer,
text_max_length=args.text_max_length,
cfg_rate=args.training_cfg_rate)
def _load_raw_data(self) -> List[Dict]:
"""Load raw data from JSON files."""
# Read folder-annotation pairs
with open(self.data_merge_path) as f:
folder_anno_pairs = [
line.strip().split(",") for line in f if line.strip()
]
assert len(
folder_anno_pairs) == 1, "Only support one folder-annotation pair"
assert len(folder_anno_pairs[0]
) == 2, "Folder-annotation pair should have two elements"
folder, annotation_file = folder_anno_pairs[0]
data_items: List[Dict] = []
with open(annotation_file) as f:
data_items = json.load(f)
# Update paths with folder prefix
for item in data_items:
item["path"] = opj(folder, item["path"])
return data_items
def _process_metadata(self) -> List[PreprocessBatch]:
"""Process the raw metadata through all filtering stages."""
raw_data = self._load_raw_data()
processed_batches = []
# Initialize counters
filter_counts = {
"validation_failed": 0,
"resolution_failed": 0,
"frame_sampling_failed": 0
}
sample_num_frames: List[int] = []
for item in raw_data:
batch = PreprocessBatch(path=item["path"],
cap=item["cap"],
resolution=item.get("resolution"),
fps=item.get("fps"),
duration=item.get("duration"))
# Apply filtering stages
if not self._apply_filter_stages(batch, filter_counts):
continue
# Apply frame sampling processing
batch = self.frame_sampling_stage.process(
batch, temporal_sample_fn=self.temporal_sample)
processed_batches.append(batch)
assert batch.sample_num_frames is not None
sample_num_frames.append(batch.sample_num_frames)
self._log_filtering_stats(filter_counts, sample_num_frames,
len(raw_data), len(processed_batches))
return processed_batches
def _apply_filter_stages(self, batch: PreprocessBatch,
filter_counts: Dict[str, int]) -> bool:
"""Apply all filter stages and update counters. Returns True if batch should be kept."""
if not self.validation_stage.should_keep(batch):
filter_counts["validation_failed"] += 1
return False
if not self.resolution_filter_stage.should_keep(batch):
filter_counts["resolution_failed"] += 1
return False
if not self.frame_sampling_stage.should_keep(batch):
filter_counts["frame_sampling_failed"] += 1
return False
return True
def _log_filtering_stats(self, filter_counts: Dict[str, int],
sample_num_frames: List[int], before_count: int,
after_count: int):
"""Log filtering statistics."""
logger.info(
"validation_failed: %d, resolution_failed: %d, frame_sampling_failed: %d, "
"Counter(sample_num_frames): %s, before filter: %d, after filter: %d",
filter_counts['validation_failed'],
filter_counts['resolution_failed'],
filter_counts['frame_sampling_failed'], Counter(sample_num_frames),
before_count, after_count)
def __iter__(self):
"""Iterate through processed data items."""
for idx in range(len(self.processed_batches)):
yield self._get_item(idx)
def __len__(self):
return len(self.processed_batches)
def _get_item(self, idx: int) -> Dict:
"""Get a single processed data item."""
batch = self.processed_batches[idx]
# Apply transformation stages
batch = self.video_transform_stage.process(batch)
batch = self.image_transform_stage.process(batch)
batch = self.text_encoding_stage.process(batch)
# Build result dictionary
result = {
"pixel_values": batch.pixel_values,
"text": batch.text,
"input_ids": batch.input_ids,
"cond_mask": batch.cond_mask,
"path": batch.path,
}
# Add video-specific fields
if batch.is_video:
result.update({"fps": batch.fps, "duration": batch.duration})
return result
def state_dict(self) -> Dict[str, Any]:
"""Return state dict for checkpointing."""
return {"processed_batches": self.processed_batches}
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
"""Load state dict from checkpoint."""
self.processed_batches = state_dict["processed_batches"]
-352
View File
@@ -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
+220
View File
@@ -0,0 +1,220 @@
import random
from typing import Any, Dict, List, cast
import numpy as np
import torch
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
"""
Pad or crop an embedding [L, D] to exactly padding_length tokens.
Return:
- [L, D] tensor in pinned CPU memory
- [L] attention mask in pinned CPU memory
"""
L, D = t.shape
if padding_length > L: # pad
pad = torch.zeros(padding_length - L, D, dtype=t.dtype, device=t.device)
return torch.cat([t, pad], 0), torch.cat(
[torch.ones(L), torch.zeros(padding_length - L)], 0)
else: # crop
return t[:padding_length], torch.ones(padding_length)
def get_torch_tensors_from_row_dict(row_dict, keys, cfg_rate) -> Dict[str, Any]:
"""
Get the latents and prompts from a row dictionary.
"""
return_dict = {}
for key in keys:
shape, bytes = None, None
if isinstance(key, tuple):
for k in key:
try:
shape = row_dict[f"{k}_shape"]
bytes = row_dict[f"{k}_bytes"]
except KeyError:
continue
key = key[0]
if shape is None or bytes is None:
raise ValueError(f"Key {key} not found in row_dict")
else:
shape = row_dict[f"{key}_shape"]
bytes = row_dict[f"{key}_bytes"]
# TODO (peiyuan): read precision
if key == 'text_embedding' and random.random() < cfg_rate:
data = np.zeros((512, 4096), dtype=np.float32)
else:
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
data = torch.from_numpy(data)
if len(data.shape) == 3:
B, L, D = data.shape
assert B == 1, "Batch size must be 1"
data = data.squeeze(0)
return_dict[key] = data
return return_dict
def collate_latents_embs_masks(
batch_to_process,
text_padding_length,
keys,
cfg_rate=0.0
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, List[str]]:
# Initialize tensors to hold padded embeddings and masks
all_latents = []
all_embs = []
all_masks = []
caption_text = []
# Process each row individually
for i, row in enumerate(batch_to_process):
# Get tensors from row
data = get_torch_tensors_from_row_dict(row, keys, cfg_rate)
latents, emb = data["vae_latent"], data["text_embedding"]
padded_emb, mask = pad(emb, text_padding_length)
# Store in batch tensors
all_latents.append(latents)
all_embs.append(padded_emb)
all_masks.append(mask)
# TODO(py): remove this once we fix preprocess
try:
caption_text.append(row["prompt"])
except KeyError:
caption_text.append(row["caption"])
# Pin memory for faster transfer to GPU
all_latents = torch.stack(all_latents)
all_embs = torch.stack(all_embs)
all_masks = torch.stack(all_masks)
return all_latents, all_embs, all_masks, caption_text
def collate_rows_from_parquet_schema(rows,
parquet_schema,
text_padding_length,
cfg_rate=0.0) -> Dict[str, Any]:
"""
Collate rows from parquet files based on the provided schema.
Dynamically processes tensor fields based on schema and returns batched data.
Args:
rows: List of row dictionaries from parquet files
parquet_schema: PyArrow schema defining the structure of the data
Returns:
Dict containing batched tensors and metadata
"""
if not rows:
return cast(Dict[str, Any], {})
# Initialize containers for different data types
batch_data: Dict[str, Any] = {}
# Get tensor and metadata field names from schema (fields ending with '_bytes')
tensor_fields = []
metadata_fields = []
for field in parquet_schema.names:
if field.endswith('_bytes'):
shape_field = field.replace('_bytes', '_shape')
dtype_field = field.replace('_bytes', '_dtype')
tensor_name = field.replace('_bytes', '')
tensor_fields.append(tensor_name)
assert shape_field in parquet_schema.names, f"Shape field {shape_field} not found in schema for field {field}. Currently we only support *_bytes fields for tensors."
assert dtype_field in parquet_schema.names, f"Dtype field {dtype_field} not found in schema for field {field}. Currently we only support *_bytes fields for tensors."
elif not field.endswith('_shape') and not field.endswith('_dtype'):
# Only add actual metadata fields, not the shape/dtype helper fields
metadata_fields.append(field)
# Process each tensor field
for tensor_name in tensor_fields:
tensor_list = []
for row in rows:
# Get tensor data from row using the existing helper function pattern
shape_key = f"{tensor_name}_shape"
bytes_key = f"{tensor_name}_bytes"
if shape_key in row and bytes_key in row:
shape = row[shape_key]
bytes_data = row[bytes_key]
if len(bytes_data) == 0:
tensor = torch.zeros(0, dtype=torch.bfloat16)
else:
# Convert bytes to tensor using float32 as default
if tensor_name == 'text_embedding' and random.random(
) < cfg_rate:
data = np.zeros((512, 4096), dtype=np.float32)
else:
data = np.frombuffer(
bytes_data, dtype=np.float32).reshape(shape).copy()
tensor = torch.from_numpy(data)
# if len(data.shape) == 3:
# B, L, D = tensor.shape
# assert B == 1, "Batch size must be 1"
# tensor = tensor.squeeze(0)
tensor_list.append(tensor)
else:
# Handle missing tensor data
tensor_list.append(torch.zeros(0, dtype=torch.bfloat16))
# Stack tensors with special handling for text embeddings
if tensor_name == 'text_embedding':
# Handle text embeddings with padding
padded_tensors = []
attention_masks = []
for tensor in tensor_list:
if tensor.numel() > 0:
padded_tensor, mask = pad(tensor, text_padding_length)
padded_tensors.append(padded_tensor)
attention_masks.append(mask)
else:
# Handle empty embeddings - assume default embedding dimension
padded_tensors.append(
torch.zeros(text_padding_length,
768,
dtype=torch.bfloat16))
attention_masks.append(torch.zeros(text_padding_length))
batch_data[tensor_name] = torch.stack(padded_tensors)
batch_data['text_attention_mask'] = torch.stack(attention_masks)
else:
# Stack all tensors to preserve batch consistency
# Don't filter out None or empty tensors as this breaks batch sizing
try:
batch_data[tensor_name] = torch.stack(tensor_list)
except ValueError as e:
shapes = [
t.shape
if t is not None and hasattr(t, 'shape') else 'None/Invalid'
for t in tensor_list
]
raise ValueError(
f"Failed to stack tensors for field '{tensor_name}'. "
f"Tensor shapes: {shapes}. "
f"All tensors in a batch must have compatible shapes. "
f"Original error: {e}") from e
# Process metadata fields into info_list
info_list = []
for row in rows:
info = {}
for field in metadata_fields:
info[field] = row.get(field, "")
# Add prompt field for backward compatibility
info["prompt"] = info.get("caption", "")
info_list.append(info)
batch_data['info_list'] = info_list
# Add caption_text for backward compatibility
if info_list and 'caption' in info_list[0]:
batch_data['caption_text'] = [info['caption'] for info in info_list]
return batch_data
+103
View File
@@ -0,0 +1,103 @@
# SPDX-License-Identifier: Apache-2.0
# adapted from: https://github.com/a-r-r-o-w/finetrainers/blob/main/finetrainers/data/dataset.py
import os
import pathlib
import datasets
import torch
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vision_utils import load_image, load_video
logger = init_logger(__name__)
class ValidationDataset(torch.utils.data.IterableDataset):
def __init__(self, filename: str):
super().__init__()
self.filename = pathlib.Path(filename)
# get directory of filename
self.dir = os.path.abspath(self.filename.parent)
if not self.filename.exists():
raise FileNotFoundError(
f"File {self.filename.as_posix()} does not exist")
if self.filename.suffix == ".csv":
data = datasets.load_dataset("csv",
data_files=self.filename.as_posix(),
split="train")
elif self.filename.suffix == ".json":
data = datasets.load_dataset("json",
data_files=self.filename.as_posix(),
split="train",
field="data")
elif self.filename.suffix == ".parquet":
data = datasets.load_dataset("parquet",
data_files=self.filename.as_posix(),
split="train")
elif self.filename.suffix == ".arrow":
data = datasets.load_dataset("arrow",
data_files=self.filename.as_posix(),
split="train")
else:
_SUPPORTED_FILE_FORMATS = [".csv", ".json", ".parquet", ".arrow"]
raise ValueError(
f"Unsupported file format {self.filename.suffix} for validation dataset. Supported formats are: {_SUPPORTED_FILE_FORMATS}"
)
self._data = data.to_iterable_dataset()
def __iter__(self):
for sample in self._data:
# For consistency reasons, we mandate that "caption" is always present in the validation dataset.
# However, since the model specifications use "prompt", we create an alias here.
sample["prompt"] = sample["caption"]
# Load image or video if the path is provided
# TODO(aryan): need to handle custom columns here for control conditions
sample["image"] = None
sample["video"] = None
if sample.get("image_path", None) is not None:
image_path = sample["image_path"]
image_path = os.path.join(self.dir, image_path)
if not pathlib.Path(image_path).is_file(
) and not image_path.startswith("http"):
logger.warning("Image file %s does not exist.", image_path)
else:
sample["image"] = load_image(image_path)
if sample.get("video_path", None) is not None:
video_path = sample["video_path"]
video_path = os.path.join(self.dir, video_path)
if not pathlib.Path(video_path).is_file(
) and not video_path.startswith("http"):
logger.warning("Video file %s does not exist.", video_path)
else:
sample["video"] = load_video(video_path)
if sample.get("control_image_path", None) is not None:
control_image_path = sample["control_image_path"]
control_image_path = os.path.join(self.dir, control_image_path)
if not pathlib.Path(control_image_path).is_file(
) and not control_image_path.startswith("http"):
logger.warning("Control Image file %s does not exist.",
control_image_path)
else:
sample["control_image"] = load_image(control_image_path)
if sample.get("control_video_path", None) is not None:
control_video_path = sample["control_video_path"]
control_video_path = os.path.join(self.dir, control_video_path)
if not pathlib.Path(control_video_path).is_file(
) and not control_video_path.startswith("http"):
logger.warning("Control Video file %s does not exist.",
control_video_path)
else:
sample["control_video"] = load_video(control_video_path)
sample = {k: v for k, v in sample.items() if v is not None}
yield sample
+5 -5
View File
@@ -3,10 +3,10 @@
from fastvideo.v1.distributed.communication_op import *
from fastvideo.v1.distributed.parallel_state import (
cleanup_dist_env_and_memory, get_dp_group, get_dp_rank, get_dp_world_size,
get_sp_group, get_sp_parallel_rank, get_sp_world_size, get_torch_device,
get_tp_group, get_tp_rank, get_tp_world_size, get_world_group,
get_world_rank, get_world_size, init_distributed_environment,
initialize_model_parallel,
get_local_torch_device, get_sp_group, get_sp_parallel_rank,
get_sp_world_size, get_tp_group, get_tp_rank, get_tp_world_size,
get_world_group, get_world_rank, get_world_size,
init_distributed_environment, initialize_model_parallel,
maybe_init_distributed_environment_and_model_parallel,
model_parallel_is_initialized)
from fastvideo.v1.distributed.utils import *
@@ -40,5 +40,5 @@ __all__ = [
"get_tp_world_size",
# Get torch device
"get_torch_device",
"get_local_torch_device",
]
+32 -6
View File
@@ -36,6 +36,7 @@ from unittest.mock import patch
import torch
import torch.distributed
import torch.distributed as dist
from torch.distributed import Backend, ProcessGroup, ReduceOp
import fastvideo.v1.envs as envs
@@ -692,6 +693,7 @@ class GroupCoordinator:
_WORLD: Optional[GroupCoordinator] = None
_NODE: Optional[GroupCoordinator] = None
def get_world_group() -> GroupCoordinator:
@@ -699,6 +701,11 @@ def get_world_group() -> GroupCoordinator:
return _WORLD
def get_node_group() -> GroupCoordinator:
assert _NODE is not None, ("node group is not initialized")
return _NODE
def init_world_group(ranks: List[int], local_rank: int,
backend: str) -> GroupCoordinator:
return GroupCoordinator(
@@ -710,6 +717,18 @@ def init_world_group(ranks: List[int], local_rank: int,
)
def init_node_group(local_rank: int, backend: str):
cpu_group = get_world_group().cpu_group
node_ranks = same_node_ranks(cpu_group)
node_size = len(node_ranks)
all_node_ranks = [
list(range(i * node_size, (i + 1) * node_size))
for i in range(dist.get_world_size() // node_size)
]
global _NODE
_NODE = init_model_parallel_group(all_node_ranks, local_rank, backend)
def init_model_parallel_group(
group_ranks: List[List[int]],
local_rank: int,
@@ -782,6 +801,8 @@ def init_distributed_environment(
else:
assert _WORLD.world_size == torch.distributed.get_world_size(), (
"world group already initialized with a different world size")
# Init a group for each node
init_node_group(local_rank, backend)
_SP: Optional[GroupCoordinator] = None
@@ -904,7 +925,7 @@ def get_dp_rank() -> int:
return get_dp_group().rank_in_group
def get_torch_device() -> torch.device:
def get_local_torch_device() -> torch.device:
"""Return the torch device for the current rank."""
return torch.device(f"cuda:{envs.LOCAL_RANK}")
@@ -1021,17 +1042,22 @@ def cleanup_dist_env_and_memory(shutdown_ray: bool = False):
"torch._C._host_emptyCache() only available in Pytorch >=2.5")
def in_the_same_node_as(pg: Union[ProcessGroup, StatelessProcessGroup],
source_rank: int = 0) -> List[bool]:
def same_node_ranks(pg: Union[ProcessGroup, StatelessProcessGroup],
source_rank: int = 0) -> List[int]:
"""
This is a collective operation that returns if each rank is in the same node
This is a collective operation that returns ranks that are in the same node
as the source rank. It tests if processes are attached to the same
memory system (shared access to shared memory).
Args:
pg: the global process group to test
source_rank: the rank to test against
Returns:
A list of ranks that are in the same node as the source rank.
"""
if isinstance(pg, ProcessGroup):
assert torch.distributed.get_backend(
pg) != torch.distributed.Backend.NCCL, (
"in_the_same_node_as should be tested with a non-NCCL group.")
"same_node_ranks should be tested with a non-NCCL group.")
# local rank inside the group
rank = torch.distributed.get_rank(group=pg)
world_size = torch.distributed.get_world_size(group=pg)
@@ -1103,7 +1129,7 @@ def in_the_same_node_as(pg: Union[ProcessGroup, StatelessProcessGroup],
rank_data = pg.broadcast_obj(is_in_the_same_node, src=i)
aggregated_data += rank_data
return [x == 1 for x in aggregated_data.tolist()]
return [i for i, x in enumerate(aggregated_data.tolist()) if x == 1]
def initialize_tensor_parallel_group(
+6 -47
View File
@@ -4,9 +4,9 @@
import argparse
import dataclasses
import os
from typing import Any, Dict, List, Optional, cast
from typing import List, cast
from fastvideo import PipelineConfig, VideoGenerator
from fastvideo import VideoGenerator
from fastvideo.v1.configs.sample.base import SamplingParam
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
from fastvideo.v1.entrypoints.cli.utils import RaiseNotImplementedAction
@@ -37,8 +37,6 @@ class GenerateSubcommand(CLISubcommand):
def cmd(self, args: argparse.Namespace) -> None:
excluded_args = ['subparser', 'config', 'dispatch_function']
FastVideoArgs.from_cli_args(args)
provided_args = {}
for k, v in vars(args).items():
if (k not in excluded_args and v is not None
@@ -66,27 +64,19 @@ class GenerateSubcommand(CLISubcommand):
init_args = {
k: v
for k, v in merged_args.items() if k in self.init_arg_names
for k, v in merged_args.items()
if k not in self.generation_arg_names
}
generation_args = {
k: v
for k, v in merged_args.items() if k in self.generation_arg_names
}
pipeline_config = PipelineConfig.from_pretrained(
merged_args['model_path'])
update_config_from_args(pipeline_config.dit_config, merged_args,
"dit_config")
update_config_from_args(pipeline_config.vae_config, merged_args,
"vae_config")
update_config_from_args(pipeline_config, merged_args)
model_path = init_args.pop('model_path')
prompt = generation_args.pop('prompt')
generator = VideoGenerator.from_pretrained(
model_path=model_path, **init_args, pipeline_config=pipeline_config)
generator = VideoGenerator.from_pretrained(model_path=model_path,
**init_args)
generator.generate_video(prompt=prompt, **generation_args)
@@ -132,34 +122,3 @@ class GenerateSubcommand(CLISubcommand):
def cmd_init() -> List[CLISubcommand]:
return [GenerateSubcommand()]
def update_config_from_args(config: Any,
args_dict: Dict[str, Any],
prefix: Optional[str] = None) -> None:
"""
Update configuration object from arguments dictionary.
Args:
config: The configuration object to update
args_dict: Dictionary containing arguments
prefix: Prefix for the configuration parameters in the args_dict.
If None, assumes direct attribute mapping without prefix.
"""
# Handle top-level attributes (no prefix)
if prefix is None:
for key, value in args_dict.items():
if hasattr(config, key) and value is not None:
if key == "text_encoder_precisions" and isinstance(value, list):
setattr(config, key, tuple(value))
else:
setattr(config, key, value)
return
# Handle nested attributes with prefix
prefix_with_dot = f"{prefix}."
for key, value in args_dict.items():
if key.startswith(prefix_with_dot) and value is not None:
attr_name = key[len(prefix_with_dot):]
if hasattr(config, attr_name):
setattr(config, attr_name, value)
+12 -34
View File
@@ -18,8 +18,6 @@ import torch
import torchvision
from einops import rearrange
from fastvideo.v1.configs.pipelines import (PipelineConfig,
get_pipeline_config_cls_for_name)
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
@@ -55,9 +53,6 @@ class VideoGenerator:
model_path: str,
device: Optional[str] = None,
torch_dtype: Optional[torch.dtype] = None,
pipeline_config: Optional[
Union[str
| PipelineConfig]] = None,
**kwargs) -> "VideoGenerator":
"""
Create a video generator from a pretrained model.
@@ -66,35 +61,17 @@ class VideoGenerator:
model_path: Path or identifier for the pretrained model
device: Device to load the model on (e.g., "cuda", "cuda:0", "cpu")
torch_dtype: Data type for model weights (e.g., torch.float16)
**kwargs: Additional arguments to customize model loading
pipeline_config: Pipeline config to use for inference
**kwargs: Additional arguments to customize model loading, set any FastVideoArgs or PipelineConfig attributes here.
Returns:
The created video generator
Priority level: Default pipeline config < User's pipeline config < User's kwargs
"""
config = None
# 1. If users provide a pipeline config, it will override the default pipeline config
if isinstance(pipeline_config, PipelineConfig):
config = pipeline_config
else:
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
if isinstance(pipeline_config, str):
config.load_from_json(pipeline_config)
# 2. If users also provide some kwargs, it will override the pipeline config.
# The user kwargs shouldn't contain model config parameters!
if config is None:
logger.warning("No config found for model %s, using default config",
model_path)
config_args = kwargs
else:
config_args = shallow_asdict(config)
config_args.update(kwargs)
fastvideo_args = FastVideoArgs(model_path=model_path, **config_args)
# If users also provide some kwargs, it will override the FastVideoArgs and PipelineConfig.
kwargs['model_path'] = model_path
fastvideo_args = FastVideoArgs.from_kwargs(kwargs)
return cls.from_fastvideo_args(fastvideo_args)
@@ -150,16 +127,17 @@ class VideoGenerator:
"""
# Create a copy of inference args to avoid modifying the original
fastvideo_args = self.fastvideo_args
pipeline_config = fastvideo_args.pipeline_config
# Validate inputs
if not isinstance(prompt, str):
raise TypeError(
f"`prompt` must be a string, but got {type(prompt)}")
prompt = prompt.strip()
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
kwargs["prompt"] = prompt
sampling_param.update(kwargs)
@@ -176,10 +154,10 @@ class VideoGenerator:
f"height={sampling_param.height}, width={sampling_param.width}, "
f"num_frames={sampling_param.num_frames}")
temporal_scale_factor = fastvideo_args.vae_config.arch_config.temporal_compression_ratio
temporal_scale_factor = pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = sampling_param.num_frames
num_gpus = fastvideo_args.num_gpus
use_temporal_scaling_frames = fastvideo_args.vae_config.use_temporal_scaling_frames
use_temporal_scaling_frames = pipeline_config.vae_config.use_temporal_scaling_frames
# Adjust number of frames based on number of GPUs
if use_temporal_scaling_frames:
@@ -238,18 +216,18 @@ class VideoGenerator:
num_videos_per_prompt: {sampling_param.num_videos_per_prompt}
guidance_scale: {sampling_param.guidance_scale}
n_tokens: {n_tokens}
flow_shift: {fastvideo_args.flow_shift}
embedded_guidance_scale: {fastvideo_args.embedded_cfg_scale}
flow_shift: {fastvideo_args.pipeline_config.flow_shift}
embedded_guidance_scale: {fastvideo_args.pipeline_config.embedded_cfg_scale}
save_video: {sampling_param.save_video}
output_path: {sampling_param.output_path}
""" # type: ignore[attr-defined]
logger.info(debug_str)
# Prepare batch
batch = ForwardBatch(
**shallow_asdict(sampling_param),
eta=0.0,
n_tokens=n_tokens,
VSA_sparsity=fastvideo_args.VSA_sparsity,
extra={},
)
+125 -239
View File
@@ -6,26 +6,32 @@ import argparse
import dataclasses
from contextlib import contextmanager
from dataclasses import field
from typing import Any, Callable, List, Optional, Tuple
from typing import Any, Dict, List, Optional
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig, STA_Mode
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import FlexibleArgumentParser, StoreBoolean
logger = init_logger(__name__)
def preprocess_text(prompt: str) -> str:
return prompt
def clean_cli_args(args: argparse.Namespace) -> Dict[str, Any]:
"""
Clean the arguments by removing the ones that not explicitly provided by the user.
"""
provided_args = {}
for k, v in vars(args).items():
if (v is not None and hasattr(args, '_provided')
and k in args._provided):
provided_args[k] = v
def postprocess_text(output: Any) -> Any:
raise NotImplementedError
return provided_args
# args for fastvideo framework
@dataclasses.dataclass
class FastVideoArgs:
# Model and path configuration
# Model and path configuration (for convenience)
model_path: str
# Cache strategy
@@ -44,70 +50,34 @@ class FastVideoArgs:
num_gpus: int = 1
tp_size: int = -1
sp_size: int = -1
dp_size: int = 1
dp_shards: int = -1
hsdp_replicate_dim: int = 1
hsdp_shard_dim: int = -1
dist_timeout: Optional[int] = None # timeout for torch.distributed
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
pipeline_config: PipelineConfig = field(default_factory=PipelineConfig)
output_type: str = "pil"
# DiT configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
precision: str = "bf16"
use_cpu_offload: bool = True
use_cpu_offload: bool = True # For DiT
use_fsdp_inference: bool = True
text_encoder_offload: bool = True
pin_cpu_memory: bool = True
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True # Might change in between forward passes
vae_sp: bool = False # Might change in between forward passes
# vae_scale_factor: Optional[int] = None # Deprecated
vae_config: VAEConfig = field(default_factory=VAEConfig)
# Image encoder configuration
image_encoder_precision: str = "fp32"
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
# Text encoder configuration
DEFAULT_TEXT_ENCODER_PRECISIONS = (
"fp16",
# "fp16",
)
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS)
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (EncoderConfig(), ))
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: Tuple[Callable[[Any], Any], ...] = field(
default_factory=lambda: (postprocess_text, ))
# STA parameters
STA_mode: Optional[str] = None
skip_time_steps: int = 15
# LoRA parameters
lora_path: Optional[str] = None
lora_nickname: Optional[
str] = "default" # for swapping adapters in the pipeline
lora_target_names: Optional[List[
str]] = None # can restrict list of layers to adapt, e.g. ["q_proj"]
# STA parameters
# STA (Sliding Tile Attention) parameters
mask_strategy_file_path: Optional[str] = None
STA_mode: STA_Mode = STA_Mode.STA_INFERENCE
skip_time_steps: int = 15
# Compilation
enable_torch_compile: bool = False
disable_autocast: bool = False
# StepVideo specific parameters
pos_magic: Optional[str] = None
neg_magic: Optional[str] = None
timesteps_scale: Optional[bool] = None
# VSA parameters
VSA_sparsity: float = 0.0 # inference/validation sparsity
# Logging
log_level: str = "info"
# Stage verification
enable_stage_verification: bool = True
@property
def training_mode(self) -> bool:
@@ -125,11 +95,6 @@ class FastVideoArgs:
help=
"The path of the model weights. This can be a local folder or a Hugging Face repo ID.",
)
parser.add_argument(
"--dit-weight",
type=str,
help="Path to the DiT model weights",
)
parser.add_argument(
"--model-dir",
type=str,
@@ -175,31 +140,27 @@ class FastVideoArgs:
help="The number of GPUs to use.",
)
parser.add_argument(
"--tensor-parallel-size",
"--tp-size",
type=int,
default=FastVideoArgs.tp_size,
help="The tensor parallelism size.",
)
parser.add_argument(
"--sequence-parallel-size",
"--sp-size",
type=int,
default=FastVideoArgs.sp_size,
help="The sequence parallelism size.",
)
parser.add_argument(
"--data-parallel-size",
"--dp-size",
"--hsdp-replicate-dim",
type=int,
default=FastVideoArgs.dp_size,
default=FastVideoArgs.hsdp_replicate_dim,
help="The data parallelism size.",
)
parser.add_argument(
"--data-parallel-shards",
"--dp-shards",
"--hsdp-shard-dim",
type=int,
default=FastVideoArgs.dp_shards,
default=FastVideoArgs.hsdp_shard_dim,
help="The data parallelism shards.",
)
parser.add_argument(
@@ -209,19 +170,7 @@ class FastVideoArgs:
help="Set timeout for torch.distributed initialization.",
)
parser.add_argument(
"--embedded-cfg-scale",
type=float,
default=FastVideoArgs.embedded_cfg_scale,
help="Embedded CFG scale",
)
parser.add_argument(
"--flow-shift",
"--shift",
type=float,
default=FastVideoArgs.flow_shift,
help="Flow shift parameter",
)
# Output type
parser.add_argument(
"--output-type",
type=str,
@@ -230,62 +179,14 @@ class FastVideoArgs:
help="Output type for the generated video",
)
parser.add_argument(
"--precision",
type=str,
default=FastVideoArgs.precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for the model",
)
# VAE configuration
parser.add_argument(
"--vae-precision",
type=str,
default=FastVideoArgs.vae_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for VAE",
)
parser.add_argument(
"--vae-tiling",
action=StoreBoolean,
default=FastVideoArgs.vae_tiling,
help="Enable VAE tiling",
)
parser.add_argument(
"--vae-sp",
action=StoreBoolean,
help="Enable VAE spatial parallelism",
)
parser.add_argument(
"--text-encoder-precisions",
nargs="+",
type=str,
default=FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS,
choices=["fp32", "fp16", "bf16"],
help="Precision for each text encoder",
)
# Image encoder config
parser.add_argument(
"--image-encoder-precision",
type=str,
default=FastVideoArgs.image_encoder_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for image encoder",
)
# STA parameters
# STA (Sliding Tile Attention) parameters
parser.add_argument(
"--STA-mode",
type=str,
default=FastVideoArgs.STA_mode,
choices=[
"STA_inference", "STA_searching", "STA_tuning",
"STA_tuning_cfg", None
],
help="STA mode",
default=FastVideoArgs.STA_mode.value,
choices=[mode.value for mode in STA_Mode],
help=
"STA mode contains STA_inference, STA_searching, STA_tuning, STA_tuning_cfg, None",
)
parser.add_argument(
"--skip-time-steps",
@@ -309,7 +210,7 @@ class FastVideoArgs:
"--use-cpu-offload",
action=StoreBoolean,
help=
"Use CPU offload for model inference. Enable if run out of memory with FSDP.",
"Use CPU offload for DiT inference. Enable if run out of memory with FSDP.",
)
parser.add_argument(
"--use-fsdp-inference",
@@ -317,7 +218,19 @@ class FastVideoArgs:
help=
"Use FSDP for inference by sharding the model weights. Latency is very low due to prefetch--enable if run out of memory.",
)
parser.add_argument(
"--text-encoder-cpu-offload",
action=StoreBoolean,
help=
"Use CPU offload for text encoder. Enable if run out of memory.",
)
parser.add_argument(
"--pin-cpu-memory",
action=StoreBoolean,
help=
"Pin memory for CPU offload. Only added as a temp workaround if it throws \"CUDA error: invalid argument\". "
"Should be enabled in almost all cases",
)
parser.add_argument(
"--disable-autocast",
action=StoreBoolean,
@@ -325,91 +238,69 @@ class FastVideoArgs:
"Disable autocast for denoising loop and vae decoding in pipeline sampling",
)
# VSA parameters
parser.add_argument(
"--pos_magic",
type=str,
default=FastVideoArgs.pos_magic,
help="Positive magic prompt for sampling",
)
parser.add_argument(
"--neg_magic",
type=str,
default=FastVideoArgs.neg_magic,
help="Negative magic prompt for sampling",
)
parser.add_argument(
"--timesteps_scale",
type=bool,
default=FastVideoArgs.timesteps_scale,
help="Bool for applying scheduler scale in set_timesteps",
"--VSA-sparsity",
type=float,
default=FastVideoArgs.VSA_sparsity,
help="Validation sparsity for VSA",
)
# Logging
# Stage verification
parser.add_argument(
"--log-level",
type=str,
default=FastVideoArgs.log_level,
help="The logging level of all loggers.",
"--enable-stage-verification",
action=StoreBoolean,
default=FastVideoArgs.enable_stage_verification,
help="Enable input/output verification for pipeline stages",
)
# Add VAE configuration arguments
from fastvideo.v1.configs.models.vaes.base import VAEConfig
VAEConfig.add_cli_args(parser)
# Add DiT configuration arguments
from fastvideo.v1.configs.models.dits.base import DiTConfig
DiTConfig.add_cli_args(parser)
# Add pipeline configuration arguments
PipelineConfig.add_cli_args(parser)
return parser
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "FastVideoArgs":
args.tp_size = args.tensor_parallel_size
args.sp_size = args.sequence_parallel_size
args.flow_shift = getattr(args, "shift", args.flow_shift)
provided_args = clean_cli_args(args)
# Get all fields from the dataclass
attrs = [attr.name for attr in dataclasses.fields(cls)]
# Create a dictionary of attribute values, with defaults for missing attributes
kwargs = {}
for attr in attrs:
# Handle renamed attributes or those with multiple CLI names
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
kwargs[attr] = args.tensor_parallel_size
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
kwargs[attr] = args.sequence_parallel_size
elif attr == 'dp_size' and hasattr(args, 'data_parallel_size'):
kwargs[attr] = args.data_parallel_size
elif attr == 'dp_shards' and hasattr(args, 'data_parallel_shards'):
kwargs[attr] = args.data_parallel_shards
elif attr == 'flow_shift' and hasattr(args, 'shift'):
kwargs[attr] = args.shift
if attr == 'pipeline_config':
pipeline_config = PipelineConfig.from_kwargs(provided_args)
kwargs[attr] = pipeline_config
# Use getattr with default value from the dataclass for potentially missing attributes
else:
default_value = getattr(cls, attr, None)
value = getattr(args, attr, default_value)
if value is not None:
kwargs[attr] = value
kwargs[attr] = value # type: ignore
return cls(**kwargs) # type: ignore
@classmethod
def from_kwargs(cls, kwargs: Dict[str, Any]) -> "FastVideoArgs":
kwargs['pipeline_config'] = PipelineConfig.from_kwargs(kwargs)
return cls(**kwargs)
def check_fastvideo_args(self) -> None:
"""Validate inference arguments for consistency"""
if not self.inference_mode:
assert self.dp_size is not -1, "dp_size must be set for training"
assert self.dp_shards is not -1, "dp_shards must be set for training"
assert self.sp_size is not -1, "sp_size must be set for training"
assert self.hsdp_replicate_dim != -1, "hsdp_replicate_dim must be set for training"
assert self.hsdp_shard_dim != -1, "hsdp_shard_dim must be set for training"
assert self.sp_size != -1, "sp_size must be set for training"
if self.tp_size is -1:
if self.tp_size == -1:
self.tp_size = self.num_gpus
if self.sp_size is -1:
if self.sp_size == -1:
self.sp_size = self.num_gpus
if self.dp_shards is -1:
self.dp_shards = self.num_gpus
if self.hsdp_shard_dim == -1:
self.hsdp_shard_dim = self.num_gpus
assert self.sp_size <= self.num_gpus and self.num_gpus % self.sp_size == 0, "num_gpus must >= and be divisible by sp_size"
assert self.dp_size <= self.num_gpus and self.num_gpus % self.dp_size == 0, "num_gpus must >= and be divisible by dp_size"
assert self.dp_shards <= self.num_gpus and self.num_gpus % self.dp_shards == 0, "num_gpus must >= and be divisible by dp_shards"
assert self.hsdp_replicate_dim <= self.num_gpus and self.num_gpus % self.hsdp_replicate_dim == 0, "num_gpus must >= and be divisible by hsdp_replicate_dim"
assert self.hsdp_shard_dim <= self.num_gpus and self.num_gpus % self.hsdp_shard_dim == 0, "num_gpus must >= and be divisible by hsdp_shard_dim"
if self.num_gpus < max(self.tp_size, self.sp_size):
self.num_gpus = max(self.tp_size, self.sp_size)
@@ -419,33 +310,17 @@ class FastVideoArgs:
f"tp_size ({self.tp_size}) must be equal to sp_size ({self.sp_size})"
)
# Validate VAE spatial parallelism with VAE tiling
if self.vae_sp and not self.vae_tiling:
raise ValueError(
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
)
if len(self.text_encoder_configs) != len(self.text_encoder_precisions):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
)
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
if len(self.preprocess_text_funcs) != len(self.postprocess_text_funcs):
raise ValueError(
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
if self.enable_torch_compile and self.num_gpus > 1:
logger.warning(
"Currently torch compile does not work with multi-gpu. Setting enable_torch_compile to False"
)
self.enable_torch_compile = False
if self.pipeline_config is None:
raise ValueError("pipeline_config is not set in FastVideoArgs")
self.pipeline_config.check_pipeline_config()
_current_fastvideo_args = None
@@ -519,21 +394,22 @@ class TrainingArgs(FastVideoArgs):
# text encoder & vae & diffusion model
pretrained_model_name_or_path: str = ""
dit_model_name_or_path: str = ""
cache_dir: str = ""
# diffusion setting
ema_decay: float = 0.0
ema_start_step: int = 0
cfg: float = 0.0
training_cfg_rate: float = 0.0
precondition_outputs: bool = False
# validation & logs
validation_prompt_dir: str = ""
validation_dataset_file: str = ""
validation_preprocessed_path: str = ""
validation_sampling_steps: str = ""
validation_guidance_scale: str = ""
validation_steps: float = 0.0
log_validation: bool = False
tracker_project_name: str = ""
wandb_run_name: str = ""
seed: Optional[int] = None
# output
@@ -541,7 +417,6 @@ class TrainingArgs(FastVideoArgs):
checkpoints_total_limit: int = 0
checkpointing_steps: int = 0
resume_from_checkpoint: bool = False
logging_dir: str = ""
# optimizer & scheduler
num_train_epochs: int = 0
@@ -582,36 +457,29 @@ class TrainingArgs(FastVideoArgs):
# master_weight_type
master_weight_type: str = ""
# For fast checking in LoRA pipeline
training_mode: bool = True
# VSA training decay parameters
VSA_decay_rate: float = 0.01 # decay rate -> 0.02
VSA_decay_interval_steps: int = 1 # decay interval steps -> 50
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
provided_args = clean_cli_args(args)
# Get all fields from the dataclass
attrs = [attr.name for attr in dataclasses.fields(cls)]
logger.info(provided_args)
# Create a dictionary of attribute values, with defaults for missing attributes
kwargs = {}
for attr in attrs:
# Handle renamed attributes or those with multiple CLI names
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
kwargs[attr] = args.tensor_parallel_size
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
kwargs[attr] = args.sequence_parallel_size
elif attr == 'flow_shift' and hasattr(args, 'shift'):
kwargs[attr] = args.shift
elif attr == 'dp_size' and hasattr(args, 'data_parallel_size'):
kwargs[attr] = args.data_parallel_size
elif attr == 'dp_shards' and hasattr(args, 'data_parallel_shards'):
kwargs[attr] = args.data_parallel_shards
if attr == 'pipeline_config':
pipeline_config = PipelineConfig.from_kwargs(provided_args)
kwargs[attr] = pipeline_config
# Use getattr with default value from the dataclass for potentially missing attributes
else:
default_value = getattr(cls, attr, None)
value = getattr(args, attr, default_value)
if value is not None:
kwargs[attr] = value
kwargs[attr] = value # type: ignore
return cls(**kwargs)
return cls(**kwargs) # type: ignore
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
@@ -674,7 +542,7 @@ class TrainingArgs(FastVideoArgs):
type=int,
default=0,
help="Step to start EMA")
parser.add_argument("--cfg",
parser.add_argument("--training-cfg-rate",
type=float,
help="Classifier-free guidance scale")
parser.add_argument(
@@ -683,9 +551,12 @@ class TrainingArgs(FastVideoArgs):
help="Whether to precondition the outputs of the model")
# Validation and logging
parser.add_argument("--validation-prompt-dir",
parser.add_argument("--validation-dataset-file",
type=str,
help="Directory containing validation prompts")
help="Path to unprocessed validation dataset")
parser.add_argument("--validation-preprocessed-path",
type=str,
help="Path to processed validation dataset")
parser.add_argument("--validation-sampling-steps",
type=str,
help="Validation sampling steps")
@@ -701,6 +572,9 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--tracker-project-name",
type=str,
help="Project name for tracking")
parser.add_argument("--wandb-run-name",
type=str,
help="Run name for wandb")
parser.add_argument("--seed",
type=int,
default=42,
@@ -839,4 +713,16 @@ class TrainingArgs(FastVideoArgs):
type=str,
help="Master weight type")
# VSA parameters for training with dense to sparse adaption
parser.add_argument(
"--VSA-decay-rate", # decay rate, how much sparsity you want to decay each step
type=float,
default=TrainingArgs.VSA_decay_rate,
help="VSA decay rate")
parser.add_argument(
"--VSA-decay-interval-steps", # how many steps for training with current sparsity
type=int,
default=TrainingArgs.VSA_decay_interval_steps,
help="VSA decay interval steps")
return parser
+9 -8
View File
@@ -5,15 +5,16 @@ import time
from collections import defaultdict
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Optional
from typing import TYPE_CHECKING, Optional
import torch
# if TYPE_CHECKING:
from fastvideo.v1.attention import AttentionMetadata
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
if TYPE_CHECKING:
from fastvideo.v1.attention import AttentionMetadata
from fastvideo.v1.pipelines import ForwardBatch
logger = init_logger(__name__)
@@ -36,13 +37,13 @@ class ForwardContext:
# attn_layers: Dict[str, Any]
# TODO: extend to support per-layer dynamic forward context
attn_metadata: "AttentionMetadata" # set dynamically for each forward pass
forward_batch: Optional[ForwardBatch] = None
forward_batch: Optional["ForwardBatch"] = None
_forward_context: Optional[ForwardContext] = None
_forward_context: Optional["ForwardContext"] = None
def get_forward_context() -> ForwardContext:
def get_forward_context() -> "ForwardContext":
"""Get the current forward context."""
assert _forward_context is not None, (
"Forward context is not set. "
@@ -54,7 +55,7 @@ def get_forward_context() -> ForwardContext:
@contextmanager
def set_forward_context(current_timestep,
attn_metadata,
forward_batch: Optional[ForwardBatch] = None,
forward_batch: Optional["ForwardBatch"] = None,
fastvideo_args: Optional[FastVideoArgs] = None):
"""A context manager that stores the current forward context,
can be attention metadata, etc.
+49 -10
View File
@@ -5,6 +5,8 @@ from typing import Optional, Tuple, Union
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.distributed.tensor import DTensor
from fastvideo.v1.layers.custom_op import CustomOp
@@ -69,7 +71,12 @@ class RMSNorm(CustomOp):
x = x * torch.rsqrt(variance + self.variance_epsilon)
x = x.to(orig_dtype)
if self.has_weight:
x = x * self.weight
# TODO(wenxuan): When using CPU offload, FSDP has a bug that doesn't unwrap DTensor in final_layer_norm.
# Report this
if isinstance(self.weight, DTensor):
x = x * self.weight.to(x.device).full_tensor()
else:
x = x * self.weight
if residual is None:
return x
else:
@@ -95,6 +102,22 @@ class ScaleResidual(nn.Module):
return residual + x * gate
# adapted from Diffusers: https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/normalization.py
# NOTE(will): Needed to match behavior of diffusers and wan2.1 even while using
# FSDP's MixedPrecisionPolicy
class FP32LayerNorm(nn.LayerNorm):
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
origin_dtype = inputs.dtype
return F.layer_norm(
inputs.float(),
self.normalized_shape,
self.weight.float() if self.weight is not None else None,
self.bias.float() if self.bias is not None else None,
self.eps,
).to(origin_dtype)
class ScaleResidualLayerNormScaleShift(nn.Module):
"""
Fused operation that combines:
@@ -112,6 +135,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
eps: float = 1e-6,
elementwise_affine: bool = False,
dtype: torch.dtype = torch.float32,
compute_dtype: torch.dtype | None = None,
prefix: str = "",
):
super().__init__()
@@ -121,10 +145,15 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
eps=eps,
dtype=dtype)
elif norm_type == "layer":
self.norm = nn.LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps,
dtype=dtype)
if compute_dtype == torch.float32:
self.norm = FP32LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps)
else:
self.norm = nn.LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps,
dtype=dtype)
else:
raise NotImplementedError(f"Norm type {norm_type} not implemented")
@@ -163,18 +192,25 @@ class LayerNormScaleShift(nn.Module):
eps: float = 1e-6,
elementwise_affine: bool = False,
dtype: torch.dtype = torch.float32,
compute_dtype: torch.dtype | None = None,
prefix: str = "",
):
super().__init__()
self.compute_dtype = compute_dtype
if norm_type == "rms":
self.norm = RMSNorm(hidden_size,
has_weight=elementwise_affine,
eps=eps)
elif norm_type == "layer":
self.norm = nn.LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps,
dtype=dtype)
if self.compute_dtype == torch.float32:
self.norm = FP32LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps)
else:
self.norm = nn.LayerNorm(hidden_size,
elementwise_affine=elementwise_affine,
eps=eps,
dtype=dtype)
else:
raise NotImplementedError(f"Norm type {norm_type} not implemented")
@@ -182,4 +218,7 @@ class LayerNormScaleShift(nn.Module):
scale: torch.Tensor) -> torch.Tensor:
"""Apply ln followed by scale and shift in a single fused operation."""
normalized = self.norm(x)
return normalized * (1.0 + scale) + shift
if self.compute_dtype == torch.float32:
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
else:
return normalized * (1.0 + scale) + shift
+2 -2
View File
@@ -114,7 +114,7 @@ def _info(logger: Logger,
if (main_process_only and is_main_process) or (local_main_process_only
and is_local_main_process):
logger.log(logging.INFO, msg, *args, **kwargs)
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
global _warned_local_main_process, _warned_main_process
@@ -134,7 +134,7 @@ def _info(logger: Logger,
_warned_main_process = True
if not main_process_only and not local_main_process_only:
logger.log(logging.INFO, msg, *args, **kwargs)
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
class _FastvideoLogger(Logger):
+6 -4
View File
@@ -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)
+9 -5
View File
@@ -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):
+17 -16
View File
@@ -26,7 +26,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.v1.layers.visual_embedding import TimestepEmbedder
from fastvideo.v1.models.dits.base import BaseDiT
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
class PatchEmbed2D(nn.Module):
@@ -139,16 +139,17 @@ class StepVideoRMSNorm(nn.Module):
class SelfAttention(nn.Module):
def __init__(self,
hidden_dim,
head_dim,
rope_split: Tuple[int, int, int] = (64, 32, 32),
bias: bool = False,
with_rope: bool = True,
with_qk_norm: bool = True,
attn_type: str = "torch",
supported_attention_backends=(_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)):
def __init__(
self,
hidden_dim,
head_dim,
rope_split: Tuple[int, int, int] = (64, 32, 32),
bias: bool = False,
with_rope: bool = True,
with_qk_norm: bool = True,
attn_type: str = "torch",
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)):
super().__init__()
self.head_dim = head_dim
self.hidden_dim = hidden_dim
@@ -257,7 +258,8 @@ class CrossAttention(nn.Module):
head_dim,
bias=False,
with_qk_norm=True,
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)
) -> None:
super().__init__()
self.head_dim = head_dim
@@ -453,11 +455,10 @@ class StepVideoTransformerBlock(nn.Module):
class StepVideoModel(BaseDiT):
# (Optional) Keep the same attribute for compatibility with splitting, etc.
_fsdp_shard_conditions = [
lambda n, m: "transformer_blocks" in n and n.split(".")[-1].isdigit(),
# lambda n, m: "pos_embed" in n # If needed for the patch embedding.
]
_fsdp_shard_conditions = StepVideoConfig()._fsdp_shard_conditions
_param_names_mapping = StepVideoConfig()._param_names_mapping
_reverse_param_names_mapping = StepVideoConfig(
)._reverse_param_names_mapping
_lora_param_names_mapping = StepVideoConfig()._lora_param_names_mapping
_supported_attention_backends = StepVideoConfig(
)._supported_attention_backends
+58 -38
View File
@@ -14,8 +14,8 @@ from fastvideo.v1.configs.models.dits import WanVideoConfig
from fastvideo.v1.configs.sample.wan import WanTeaCacheParams
from fastvideo.v1.distributed.parallel_state import get_sp_world_size
from fastvideo.v1.forward_context import get_forward_context
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, RMSNorm,
ScaleResidual,
from fastvideo.v1.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.v1.layers.linear import ReplicatedLinear
# from torch.nn import RMSNorm
@@ -26,7 +26,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
PatchEmbed, TimestepEmbedder)
from fastvideo.v1.models.dits.base import CachableDiT
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
class WanImageEmbedding(torch.nn.Module):
@@ -34,9 +34,9 @@ class WanImageEmbedding(torch.nn.Module):
def __init__(self, in_features: int, out_features: int):
super().__init__()
self.norm1 = nn.LayerNorm(in_features)
self.norm1 = FP32LayerNorm(in_features)
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
self.norm2 = nn.LayerNorm(out_features)
self.norm2 = FP32LayerNorm(out_features)
def forward(self,
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
@@ -125,8 +125,8 @@ class WanSelfAttention(nn.Module):
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=(_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA))
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA))
def forward(self, x: torch.Tensor, context: torch.Tensor,
context_lens: int):
@@ -174,7 +174,8 @@ class WanI2VCrossAttention(WanSelfAttention):
window_size=(-1, -1),
qk_norm=True,
eps=1e-6,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None
) -> None:
super().__init__(dim, num_heads, window_size, qk_norm, eps,
supported_attention_backends)
@@ -216,21 +217,22 @@ class WanI2VCrossAttention(WanSelfAttention):
class WanTransformerBlock(nn.Module):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
prefix: str = ""):
def __init__(
self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
prefix: str = ""):
super().__init__()
# 1. Self-attention
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -261,7 +263,8 @@ class WanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
@@ -281,7 +284,8 @@ class WanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -358,21 +362,22 @@ class WanTransformerBlock(nn.Module):
class WanTransformerBlock_VSA(nn.Module):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
prefix: str = ""):
def __init__(
self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
prefix: str = ""):
super().__init__()
# 1. Self-attention
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -404,7 +409,8 @@ class WanTransformerBlock_VSA(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
@@ -424,7 +430,8 @@ class WanTransformerBlock_VSA(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -511,6 +518,7 @@ class WanTransformer3DModel(CachableDiT):
_supported_attention_backends = WanVideoConfig(
)._supported_attention_backends
_param_names_mapping = WanVideoConfig()._param_names_mapping
_reverse_param_names_mapping = WanVideoConfig()._reverse_param_names_mapping
_lora_param_names_mapping = WanVideoConfig()._lora_param_names_mapping
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
@@ -561,7 +569,8 @@ class WanTransformer3DModel(CachableDiT):
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32)
dtype=torch.float32,
compute_dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
@@ -569,6 +578,17 @@ class WanTransformer3DModel(CachableDiT):
self.gradient_checkpointing = False
# For type checking
self.previous_e0_even = None
self.previous_e0_odd = None
self.previous_residual_even = None
self.previous_residual_odd = None
self.is_even = True
self.should_calc_even = True
self.should_calc_odd = True
self.accumulated_rel_l1_distance_even = 0
self.accumulated_rel_l1_distance_odd = 0
self.cnt = 0
self.__post_init__()
def forward(self,
@@ -655,7 +675,7 @@ class WanTransformer3DModel(CachableDiT):
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
hidden_states = self.norm_out(hidden_states.float(), shift, scale)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
+14 -6
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from typing import Optional, Tuple
from dataclasses import field
from typing import List, Optional, Tuple
import torch
from torch import nn
@@ -8,16 +9,22 @@ from torch import nn
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
ImageEncoderConfig,
TextEncoderConfig)
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
class TextEncoder(nn.Module, ABC):
_fsdp_shard_conditions: list = field(default_factory=lambda: [])
_stacked_params_mapping: List[Tuple[str, str,
str]] = field(default_factory=list)
_supported_attention_backends: Tuple[
_Backend, ...] = TextEncoderConfig()._supported_attention_backends
AttentionBackendEnum,
...] = TextEncoderConfig()._supported_attention_backends
def __init__(self, config: TextEncoderConfig) -> None:
super().__init__()
self.config = config
self._fsdp_shard_conditions = config._fsdp_shard_conditions
self._stacked_params_mapping = config.arch_config.stacked_params_mapping
if not self.supported_attention_backends:
raise ValueError(
f"Subclass {self.__class__.__name__} must define _supported_attention_backends"
@@ -34,13 +41,14 @@ class TextEncoder(nn.Module, ABC):
pass
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
return self._supported_attention_backends
class ImageEncoder(nn.Module, ABC):
_supported_attention_backends: Tuple[
_Backend, ...] = ImageEncoderConfig()._supported_attention_backends
AttentionBackendEnum,
...] = ImageEncoderConfig()._supported_attention_backends
def __init__(self, config: ImageEncoderConfig) -> None:
super().__init__()
@@ -56,5 +64,5 @@ class ImageEncoder(nn.Module, ABC):
pass
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
return self._supported_attention_backends
+3 -7
View File
@@ -596,12 +596,7 @@ class CLIPVisionModel(ImageEncoder):
# ref: https://github.com/vllm-project/vllm/pull/7186#discussion_r1734163986
def load_weights(self, weights: Iterable[Tuple[str,
torch.Tensor]]) -> Set[str]:
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
("qkv_proj", "q_proj", "q"),
("qkv_proj", "k_proj", "k"),
("qkv_proj", "v_proj", "v"),
]
params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set()
layer_count = len(self.vision_model.encoder.layers)
@@ -620,7 +615,8 @@ class CLIPVisionModel(ImageEncoder):
if layer_idx >= layer_count:
continue
for (param_name, weight_name, shard_id) in stacked_params_mapping:
for (param_name, weight_name,
shard_id) in self.config.arch_config.stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
+2 -9
View File
@@ -369,14 +369,7 @@ class LlamaModel(TextEncoder):
def load_weights(self, weights: Iterable[Tuple[str,
torch.Tensor]]) -> Set[str]:
stacked_params_mapping = [
# (param_name, shard_name, shard_id)
(".qkv_proj", ".q_proj", "q"),
(".qkv_proj", ".k_proj", "k"),
(".qkv_proj", ".v_proj", "v"),
(".gate_up_proj", ".gate_proj", 0),
(".gate_up_proj", ".up_proj", 1),
]
params_dict = dict(self.named_parameters())
loaded_params: Set[str] = set()
for name, loaded_weight in weights:
@@ -406,7 +399,7 @@ class LlamaModel(TextEncoder):
continue
else:
name = kv_scale_name
for param_name, weight_name, shard_id in stacked_params_mapping:
for param_name, weight_name, shard_id in self.config.arch_config.stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)

Some files were not shown because too many files have changed in this diff Show More