Compare commits

..
Author SHA1 Message Date
SolitaryThinker 48f9690c47 load_video 2025-06-24 10:46:20 -07:00
SolitaryThinker 32df3bcca3 update i2v script 2025-06-24 03:20:09 -07:00
SolitaryThinker b8c81191b6 improve script format 2025-06-24 03:01:55 -07:00
SolitaryThinker e506c4074f remove print and enable first val 2025-06-24 01:55:05 -07:00
SolitaryThinker 1b3914da84 update scripts 2025-06-22 21:09:30 -07:00
SolitaryThinker f471fd3f02 i2v working 2025-06-22 20:59:37 -07:00
SolitaryThinker 2e6d5c5304 t2v working again 2025-06-22 19:45:44 -07:00
SolitaryThinker 5db34184c7 t2v example 2025-06-22 21:53:05 +00:00
SolitaryThinker 37252bf62c f 2025-06-22 11:58:18 +00:00
SolitaryThinker e9263f7d2b update 2025-06-22 04:20:34 -07:00
SolitaryThinker 93afb86c20 update 2025-06-22 04:19:10 -07:00
SolitaryThinker 0d944ba9c1 slrm 2025-06-22 04:09:09 -07:00
SolitaryThinker 3d78604281 update path 2025-06-22 03:54:39 -07:00
SolitaryThinker 0694b0c5eb exmaple scripts 2025-06-22 03:44:57 -07:00
SolitaryThinker b6c5644d40 cleanup 2025-06-22 03:07:37 -07:00
SolitaryThinker a9089fa358 fix pil image 2025-06-22 02:29:51 -07:00
SolitaryThinker 6079a98fd7 i2v preprocess 2025-06-21 19:18:39 -07:00
SolitaryThinker 65c0fcb633 checkpoint 2025-06-21 18:02:43 -07:00
SolitaryThinker 82e3641264 checkpoint 2025-06-21 18:00:29 -07:00
92 changed files with 681 additions and 1208 deletions
+13 -95
View File
@@ -2,26 +2,25 @@ 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
- block: "Start Build"
blocked_state: "running"
prompt: "Approve build?"
- label: "Trigger Tests"
command: |
echo "Current working directory: $(pwd)"
echo "Current branch:"
git branch --show-current
echo "Full diff:"
git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD
plugins:
- monorepo-diff#v1.4.0:
diff: "git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD"
watch:
- path:
- "fastvideo/v1/models/encoders/**"
- "fastvideo/v1/models/loader/**"
- "fastvideo/v1/models/loaders/**"
- "fastvideo/v1/tests/encoders/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Encoder Tests"
@@ -32,10 +31,8 @@ steps:
queue: "default"
- path:
- "fastvideo/v1/models/vaes/**"
- "fastvideo/v1/models/loader/**"
- "fastvideo/v1/models/loaders/**"
- "fastvideo/v1/tests/vaes/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "VAE Tests"
@@ -46,12 +43,10 @@ steps:
queue: "default"
- path:
- "fastvideo/v1/models/dits/**"
- "fastvideo/v1/models/loader/**"
- "fastvideo/v1/models/loaders/**"
- "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"
@@ -60,8 +55,7 @@ steps:
- TEST_TYPE=transformer
agents:
queue: "default"
- path:
- "fastvideo/v1/**/*.py"
- path: "fastvideo/v1/**/*.py"
config:
command: "timeout 60m .buildkite/scripts/pr_test.sh"
label: "SSIM Tests"
@@ -70,79 +64,3 @@ steps:
- TEST_TYPE=ssim
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Training Tests"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=training
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "csrc/attn/vsa/**"
- "csrc/attn/tk/**"
- "csrc/attn/setup_vsa.py"
- "csrc/attn/config_vsa.py"
- "csrc/attn/vsa.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Training Tests VSA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=training_vsa
agents:
queue: "default"
- path:
- "fastvideo/v1/**"
- "csrc/attn/st_attn/**"
- "csrc/attn/setup_sta.py"
- "csrc/attn/config_sta.py"
- "csrc/attn/st_attn.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Inference Tests STA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=inference_sta
agents:
queue: "default"
- path:
- "csrc/attn/st_attn/**"
- "csrc/attn/setup_sta.py"
- "csrc/attn/config_sta.py"
- "csrc/attn/st_attn.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Precision Tests STA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=precision_sta
agents:
queue: "default"
- path:
- "csrc/attn/vsa/**"
- "csrc/attn/tk/**"
- "csrc/attn/setup_vsa.py"
- "csrc/attn/config_vsa.py"
- "csrc/attn/vsa.cpp"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 30m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VSA"
env:
- BUILDKITE_CLEAN_CHECKOUT=true
- TEST_TYPE=precision_vsa
agents:
queue: "default"
+4 -30
View File
@@ -31,10 +31,6 @@ 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
@@ -58,44 +54,22 @@ if [ -z "${TEST_TYPE:-}" ]; then
fi
log "Test type: $TEST_TYPE"
MODAL_ENV="BUILDKITE_REPO=$BUILDKITE_REPO BUILDKITE_COMMIT=$BUILDKITE_COMMIT IMAGE_VERSION=$IMAGE_VERSION"
case "$TEST_TYPE" in
"encoder")
log "Running encoder tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
MODAL_COMMAND="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"
MODAL_COMMAND="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"
MODAL_COMMAND="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"
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
;;
*)
log "Error: Unknown test type: $TEST_TYPE"
-40
View File
@@ -1,40 +0,0 @@
#!/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
+33 -103
View File
@@ -14,9 +14,13 @@ on:
- ".github/workflows/pr-test.yml"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
- "csrc/**"
workflow_dispatch:
inputs:
custom_image:
description: "Custom image from this repository (default: fastvideo-dev:py3.12-latest)"
required: false
default: "fastvideo-dev:py3.12-latest"
type: string
run_encoder_test:
description: "Run encoder-test"
required: false
@@ -52,16 +56,6 @@ on:
required: false
default: false
type: boolean
run_precision_test_STA:
description: "Run precision-test-STA"
required: false
default: false
type: boolean
run_precision_test_VSA:
description: "Run precision-test-VSA"
required: false
default: false
type: boolean
run_nightly_test:
description: "Run nightly-test"
required: false
@@ -71,7 +65,6 @@ on:
env:
PYTHONUNBUFFERED: "1"
concurrency:
group: pr-test-${{ github.ref }}
cancel-in-progress: true
@@ -91,69 +84,44 @@ jobs:
training-test: ${{ steps.filter.outputs.training-test }}
training-test-VSA: ${{ steps.filter.outputs.training-test-VSA }}
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
steps:
- uses: actions/checkout@v4
- uses: dorny/paths-filter@v3
id: filter
with:
filters: |
# Define reusable path patterns
common-paths: &common-paths
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
sta-kernel-paths: &sta-kernel-paths
- 'csrc/attn/st_attn/**'
- 'csrc/attn/setup_sta.py'
- 'csrc/attn/config_sta.py'
- 'csrc/attn/st_attn.cpp'
vsa-kernel-paths: &vsa-kernel-paths
- 'csrc/attn/vsa/**'
- 'csrc/attn/tk/**'
- 'csrc/attn/setup_vsa.py'
- 'csrc/attn/config_vsa.py'
- 'csrc/attn/vsa.cpp'
vsa-paths: &vsa-paths
- 'fastvideo/v1/**'
- *common-paths
- *vsa-kernel-paths
# Actual tests
encoder-test:
- 'fastvideo/v1/models/encoders/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/encoders/**'
- *common-paths
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
vae-test:
- 'fastvideo/v1/models/vaes/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/vaes/**'
- *common-paths
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
transformer-test:
- 'fastvideo/v1/models/dits/**'
- 'fastvideo/v1/models/loaders/**'
- 'fastvideo/v1/tests/transformers/**'
- 'fastvideo/v1/layers/**'
- 'fastvideo/v1/attention/**'
- *common-paths
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
training-test:
- 'fastvideo/v1/**'
- *common-paths
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
training-test-VSA:
- 'fastvideo/v1/**'
- *common-paths
- *vsa-kernel-paths
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
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
- 'pyproject.toml'
- 'docker/Dockerfile.python3.12'
encoder-test:
needs: change-filter
@@ -166,7 +134,7 @@ jobs:
gpu_type: "NVIDIA A40"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
timeout_minutes: 30
secrets:
@@ -184,7 +152,7 @@ jobs:
gpu_type: "NVIDIA A40"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
timeout_minutes: 30
secrets:
@@ -202,7 +170,7 @@ jobs:
gpu_type: "NVIDIA L40S"
gpu_count: 1
volume_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:py3.12-latest' }}"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
timeout_minutes: 30
secrets:
@@ -212,7 +180,8 @@ jobs:
ssim-test:
needs: change-filter
if: >-
github.event_name != 'workflow_dispatch' || (github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_ssim_test == 'true')
strategy:
fail-fast: false
matrix:
@@ -238,7 +207,7 @@ jobs:
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.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
@@ -247,7 +216,7 @@ jobs:
gpu_count: 4
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || '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:
@@ -258,16 +227,16 @@ jobs:
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.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test_VSA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "training-test-VSA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 2
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || '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:
@@ -278,60 +247,22 @@ jobs:
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.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_inference_test_STA == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "inference-test-STA"
gpu_type: "NVIDIA H100 NVL"
gpu_count: 2
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || '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')
@@ -342,7 +273,7 @@ jobs:
gpu_count: 4
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || '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:
@@ -351,8 +282,7 @@ jobs:
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
runpod-cleanup:
# 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]
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
if: ${{ always() && ((github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) || github.event_name == 'workflow_dispatch') }}
runs-on: ubuntu-latest
steps:
@@ -369,7 +299,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", "training-test", "training-test-VSA", "inference-test-STA", "precision-test-STA", "precision-test-VSA"]'
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12"]'
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
run: python .github/scripts/runpod_cleanup.py
+2 -8
View File
@@ -4,7 +4,7 @@
## Installation
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.
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
First, install C++20 for ThunderKittens:
```bash
sudo apt update
@@ -53,14 +53,8 @@ out = sliding_tile_attention(q, k, v, window_size, 0, False)
## Test
```bash
python tests/test_sta.py # test STA
python tests/test_block_sparse.py # test VSA
python test/test_sta.py
```
## 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,7 +5,6 @@ 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"):
@@ -14,16 +13,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 compute_TFLOPS(flops, ms):
flops = flops / 1e12
ms = ms / 1e3
return flops / ms
def efficiency(flop, time):
flop = flop / 1e12
time = time / 1e6
return flop / time
def benchmark_attention(configurations):
results = {'fwd': defaultdict(list), 'bwd': defaultdict(list)}
for B, H, N, D, causal, dit_seq_shape, window_size in configurations:
for B, H, N, D, causal in configurations:
print("=" * 60)
print(f"Timing forward and backward pass for B={B}, H={H}, N={N}, D={D}, causal={causal}")
@@ -31,31 +30,38 @@ 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()
# 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()
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()
# # Warmup for forward pass
# for _ in range(10):
# o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, dit_seq_shape)
# 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)]
# # 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))
torch.cuda.empty_cache()
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
# Warmup for forward pass
for _ in range(10):
o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, '18x48x80')
tflops_fwd = compute_TFLOPS(flops(B, N, H, D, causal, 'fwd'), ms)
# 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)
results['fwd'][(D, causal)].append((N, tflops_fwd))
print(f"Average time for forward pass (ms): {ms:.2f}")
print(f"Average TFLOPS: {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("-" * 60)
# torch.cuda.empty_cache()
@@ -79,14 +85,15 @@ 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 = compute_TFLOPS(flops(B, N, H, D, causal, 'bwd'), ms)
# tflops_bwd = efficiency(flops(B, N, H, D, causal, 'bwd'), time_us_bwd)
# results['bwd'][(D, causal)].append((N, tflops_bwd))
# print(f"Average time for backward pass(ms): {ms:.2f}")
# print(f"Average TFLOPS: {tflops_bwd}")
# print("=" * 60)
# 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)
torch.cuda.empty_cache()
torch.cuda.synchronize()
return results
@@ -117,10 +124,7 @@ def plot_results(results):
# Example list of configurations to test
configurations = [
(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]),
(2, 24, 69120, 128, False),
# (16, 16, 768*16, 128, False),
# (16, 16, 768*2, 128, False),
# (16, 16, 768*4, 128, False),
+22 -31
View File
@@ -4,17 +4,9 @@
#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);
@@ -125,16 +117,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_int(qt, DT, CT-DT-1);
qh = clamp_int(qh, DH, CH-DH-1);
qw = clamp_int(qw, DW, CW-DW-1);
qt = CLAMP(qt, DT, CT-DT-1);
qh = CLAMP(qh, DH, CH-DH-1);
qw = CLAMP(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_int(qt - kt) <= DT) && (abs_int(qh - kh) <= DH) && (abs_int(qw - kw) <= DW);
bool mask = (ABS(qt - kt) <= DT) && (ABS(qh - kh) <= DH) && (ABS(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));
@@ -175,15 +167,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_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);
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);
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++) {
@@ -242,7 +234,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_int(DT*2+1, 1, CT) * clamp_int(DH*2+1, 1, CH) * clamp_int(DW*2+1, 1, CW) * 3 - 1 ;
kv_iters = CLAMP(DT*2+1, 1, CT) * CLAMP(DH*2+1, 1, CH) * CLAMP(DW*2+1, 1, CW) * 3 - 1 ;
}
kittens::wait(qsmem_semaphore, 0);
@@ -423,9 +415,8 @@ 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();
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
cudaDeviceSynchronize();
auto stream = at::cuda::getCurrentCUDAStream().stream();
if (head_dim == 128) {
@@ -451,8 +442,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)};
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
int threads = NUM_WORKERS * kittens::WARP_THREADS;
auto mem_size = kittens::MAX_SHARED_MEMORY;
auto 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);
@@ -832,10 +823,10 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
}
CHECK_CUDA_ERROR(cudaGetLastError());
// cudaStreamSynchronize(stream);
cudaStreamSynchronize(stream);
}
return o;
//cudadevicesynchronize();
cudaDeviceSynchronize();
}
@@ -7,7 +7,6 @@ from vsa import BLOCK_M, BLOCK_N
import numpy as np
import random
import gc
def set_seed(seed: int = 42):
# Python random module
@@ -21,6 +20,15 @@ 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):
@@ -127,7 +135,9 @@ 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):
def main():
args = parse_arguments()
set_seed(42)
# Extract parameters
@@ -181,36 +191,23 @@ def main(args):
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
# 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()
# testing forward
o = BlockSparseAttentionFunction.apply(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
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)
@@ -218,72 +215,52 @@ def main(args):
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}, 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(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("\nGradient Q metrics:")
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(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("\nGradient K metrics:")
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(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("\nGradient V metrics:")
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}")
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}")
if __name__ == "__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)
main()
@@ -81,7 +81,5 @@ 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']}")
+19 -23
View File
@@ -3,8 +3,6 @@
#include "kittens.cuh"
#include <cooperative_groups.h>
#include <iostream>
#include <c10/cuda/CUDAGuard.h>
using namespace kittens;
namespace cg = cooperative_groups;
@@ -942,9 +940,8 @@ 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();
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
cudaDeviceSynchronize();
auto 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>;
@@ -969,7 +966,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())};
constexpr int mem_size = 54000;
auto mem_size = 54000;
dim3 grid(seq_len/(64), qo_heads, batch);
@@ -982,7 +979,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) {
@@ -1008,7 +1005,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())};
constexpr int mem_size = 54000;
auto mem_size = 54000;
dim3 grid(seq_len/(64), qo_heads, batch);
@@ -1021,11 +1018,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>
@@ -1135,14 +1132,13 @@ block_sparse_attention_backward(torch::Tensor q,
float* d_kg = reinterpret_cast<float*>(kg_ptr);
float* d_vg = reinterpret_cast<float*>(vg_ptr);
constexpr int mem_size = kittens::MAX_SHARED_MEMORY;
int threads = 4 * kittens::WARP_THREADS;
auto mem_size = kittens::MAX_SHARED_MEMORY;
auto threads = 4 * kittens::WARP_THREADS;
//cudadevicesynchronize();
const c10::cuda::OptionalCUDAGuard device_guard(q.device());
const cudaStream_t stream = at::cuda::getCurrentCUDAStream().stream();
cudaDeviceSynchronize();
auto 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);
@@ -1226,7 +1222,7 @@ block_sparse_attention_backward(torch::Tensor q,
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
threads = 128;
//cudadevicesynchronize();
cudaDeviceSynchronize();
{
cudaFuncSetAttribute(
@@ -1244,8 +1240,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;
@@ -1330,7 +1326,7 @@ block_sparse_attention_backward(torch::Tensor q,
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
threads = 128;
//cudadevicesynchronize();
cudaDeviceSynchronize();
{
cudaFuncSetAttribute(
@@ -1342,10 +1338,10 @@ block_sparse_attention_backward(torch::Tensor q,
bwd_attend_ker<128><<<grid_bwd_2, threads, 113000, stream>>>(bwd_global);
}
// cudaStreamSynchronize(stream);
//cudadevicesynchronize();
cudaStreamSynchronize(stream);
cudaDeviceSynchronize();
}
return {qg, kg, vg};
//cudadevicesynchronize();
cudaDeviceSynchronize();
}
@@ -1,10 +0,0 @@
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`
@@ -1,3 +0,0 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -1,10 +1,7 @@
#!/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
@@ -15,7 +12,7 @@ NUM_GPUS=8
training_args=(
--tracker_project_name "wan_i2v_finetune"
--output_dir "$DATA_DIR/outputs/wan_i2v_finetune"
--max_train_steps 2000
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
@@ -36,8 +33,8 @@ parallel_args=(
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
--model_path "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
--pretrained_model_name_or_path "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
)
# Dataset arguments
@@ -59,7 +56,7 @@ validation_args=(
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 1000
--checkpointing_steps 6000
--weight_decay 1e-4
--max_grad_norm 1.0
)
@@ -69,7 +66,7 @@ miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--cfg 0.0
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
@@ -1,5 +1,5 @@
#!/bin/bash
#SBATCH --job-name=i2v
#SBATCH --job-name=FV_2N_14B
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=4
@@ -7,10 +7,10 @@
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
#SBATCH --nodelist=fs-mbz-gpu-[400-550]
#SBATCH --mem=1440G
#SBATCH --output=i2v_output/i2v_%j.out
#SBATCH --error=i2v_output/i2v_%j.err
#SBATCH --output=4n_i2v/4n_i2v_%j.out
#SBATCH --error=4n_i2v/4n_i2v_%j.err
#SBATCH --exclusive
set -e -x
@@ -30,9 +30,7 @@ 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
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
@@ -40,91 +38,60 @@ 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 WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
DATA_DIR=data/crush-smol_processed_i2v/combined_parquet_dataset
VALIDATION_DIR=data/crush-smol_processed_i2v/validation_parquet_dataset
NUM_GPUS=8
# 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[@]}"
fastvideo/v1/training/wan_i2v_training_pipeline.py\
--model_path Wan-AI/Wan2.1-I2V-14B-480P-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-I2V-14B-480P-Diffusers \
--cache_dir "/home/ray/.cache"\
--data_path "$DATA_DIR"\
--validation_preprocessed_path "$VALIDATION_DIR"\
--train_batch_size=1\
--num_latent_t 16 \
--num_gpus $NUM_GPUS \
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES \
--hsdp_shard_dim $NUM_GPUS \
--train_sp_batch_size 1\
--dataloader_num_workers 10\
--gradient_accumulation_steps=2\
--max_train_steps=10000 \
--learning_rate=5e-5\
--mixed_precision="bf16"\
--checkpointing_steps=11000 \
--validation_steps 100\
--validation_sampling_steps "40" \
--log_validation \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune_2n"\
--tracker_project_name wan_i2v_finetune \
--num_height 480 \
--num_width 832 \
--num_frames 77 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--weight_decay 1e-4 \
--not_apply_cfg_solver \
--dit_precision "fp32" \
--max_grad_norm 1.0
@@ -1,5 +1,4 @@
#!/bin/bash
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
MODEL_TYPE="wan"
@@ -1,10 +0,0 @@
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`
@@ -1,3 +0,0 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -1,5 +1,3 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
@@ -18,7 +16,7 @@ training_args=(
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 8
--gradient_accumulation_steps 1
--num_latent_t 8
--num_height 480
--num_width 832
@@ -27,11 +25,11 @@ training_args=(
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size $NUM_GPUS
--tp_size $NUM_GPUS
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
--num_gpus $NUM_GPUS \
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--hsdp_replicate_dim 1 \
--hsdp_shard_dim $NUM_GPUS \
)
# Model arguments
@@ -50,17 +48,17 @@ dataset_args=(
validation_args=(
--log_validation
--validation_preprocessed_path $VALIDATION_DIR
--validation_steps 50
--validation_steps 100
--validation_sampling_steps "50"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-5
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 6000
--weight_decay 1e-4
--weight_decay 0.01
--max_grad_norm 1.0
)
@@ -69,7 +67,7 @@ miscellaneous_args=(
--inference_mode False
--allow_tf32
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--cfg 0.0
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
@@ -1,5 +1,5 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --job-name=FV_2N_14B
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=1
@@ -7,10 +7,10 @@
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --nodelist=fs-mbz-gpu-[100-850]
#SBATCH --nodelist=fs-mbz-gpu-[400-550]
#SBATCH --mem=1440G
#SBATCH --output=t2v_output/t2v_%j.out
#SBATCH --error=t2v_output/t2v_%j.err
#SBATCH --output=4n_i2v/4n_i2v_%j.out
#SBATCH --error=4n_i2v/4n_i2v_%j.err
#SBATCH --exclusive
set -e -x
@@ -30,98 +30,69 @@ 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
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=8
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DIR="data/crush-smol_processed_t2v/validation_parquet_dataset/"
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# 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
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
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[@]}"
fastvideo/v1/training/wan_training_pipeline.py\
--model_path $MODEL_PATH \
--inference_mode False\
--pretrained_model_name_or_path $MODEL_PATH \
--cache_dir "/home/ray/.cache"\
--data_path "$DATA_DIR"\
--validation_preprocessed_path "$VALIDATION_DIR"\
--train_batch_size=1\
--num_latent_t 8 \
--num_gpus $NUM_GPUS \
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES \
--hsdp_shard_dim $NUM_GPUS \
--train_sp_batch_size 1\
--dataloader_num_workers 10\
--gradient_accumulation_steps=1\
--max_train_steps=10000 \
--learning_rate=5e-5\
--mixed_precision="bf16"\
--checkpointing_steps=11000 \
--validation_steps 100\
--validation_sampling_steps "40" \
--log_validation \
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--cfg 0.0\
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune_2n"\
--tracker_project_name wan_i2v_finetune \
--num_height 480 \
--num_width 832 \
--num_frames 77 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--weight_decay 1e-4 \
--not_apply_cfg_solver \
--dit_precision "fp32" \
--max_grad_norm 1.0
@@ -0,0 +1,13 @@
{
"data": [
{
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
"image_path": null,
"video_path": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-034.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -1,5 +1,4 @@
#!/bin/bash
# export WANDB_MODE="offline"
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
-1
View File
@@ -12,7 +12,6 @@ class DiTArchConfig(ArchConfig):
_fsdp_shard_conditions: list = field(default_factory=list)
_compile_conditions: list = field(default_factory=list)
_param_names_mapping: dict = field(default_factory=dict)
_reverse_param_names_mapping: dict = field(default_factory=dict)
_lora_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
@@ -147,9 +147,6 @@ 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
+1 -5
View File
@@ -49,13 +49,9 @@ 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(
@@ -8,7 +8,6 @@ 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
@@ -68,18 +67,14 @@ def main() -> None:
# Create DataLoader with proper settings
dataset, dataloader = build_parquet_map_style_dataloader(
args.path,
args.batch_size,
parquet_schema=pyarrow_schema_t2v,
num_data_workers=args.num_data_workers)
args.path, args.batch_size, args.num_data_workers)
logger.info("Initialized dataloader with %d batches", len(dataloader))
if args.verify_resume:
# First pass - record latent sums
first_pass_sums = []
for i, batch in enumerate(dataloader):
latents = batch['vae_latent']
embeddings = batch['text_embedding']
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)
@@ -105,18 +100,14 @@ def main() -> None:
# 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)
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)
for i, batch in enumerate(dataloader):
latents = batch['vae_latent']
embeddings = batch['text_embedding']
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",
@@ -125,16 +116,11 @@ def main() -> None:
break
dataset, dataloader = build_parquet_map_style_dataloader(
args.path,
args.batch_size,
parquet_schema=pyarrow_schema_t2v,
num_data_workers=args.num_data_workers)
args.path, args.batch_size, args.num_data_workers)
# Second pass - verify latent sums match
second_pass_sums = []
for i, batch in enumerate(dataloader):
latents = batch['vae_latent']
embeddings = batch['text_embedding']
for i, (latents, embeddings, masks) 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)",
@@ -158,9 +144,8 @@ def main() -> None:
total_samples = 0
total_batches = 0
for _ in range(args.num_epoch):
for i, batch in enumerate(dataloader):
latents = batch['vae_latent']
embeddings = batch['text_embedding']
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
if i >= args.num_batches_per_epoch:
break
@@ -185,6 +185,9 @@ class LatentsParquetMapStyleDataset(Dataset):
Note:
Using parquet for map style dataset is not efficient, we mainly keep it for backward compatibility and debugging.
"""
# Modify this in the future if we want to add more keys, for example, in image to video.
keys = [("vae_latent", "latent"), "text_embedding", "clip_feature",
"first_frame_latent", "pil_image"]
def __init__(
self,
@@ -201,6 +204,10 @@ class LatentsParquetMapStyleDataset(Dataset):
self.path = path
self.cfg_rate = cfg_rate
self.parquet_schema = parquet_schema
if cfg_rate > 0.0:
raise ValueError(
"cfg_rate > 0.0 is not supported for now because it will trigger bug when num_data_workers > 0"
)
logger.info("Initializing LatentsParquetMapStyleDataset with path: %s",
path)
self.parquet_files, self.lengths = get_parquet_files_and_length(path)
@@ -236,8 +243,7 @@ class LatentsParquetMapStyleDataset(Dataset):
batch = collate_rows_from_parquet_schema([row_dict],
self.parquet_schema,
self.text_padding_length,
cfg_rate=0.0)
self.text_padding_length)
negative_prompt = batch['info_list'][0]['prompt']
negative_prompt_embedding = batch['text_embedding']
negative_prompt_attention_mask = batch['text_attention_mask']
@@ -259,10 +265,11 @@ class LatentsParquetMapStyleDataset(Dataset):
for idx in indices
]
batch = collate_rows_from_parquet_schema(rows,
self.parquet_schema,
self.text_padding_length,
cfg_rate=self.cfg_rate)
# all_latents, all_embs, all_masks, caption_text, all_extra_latents, all_infos = collate_latents_embs_masks(
# rows, self.text_padding_length, self.keys)
# return all_latents, all_embs, all_masks, caption_text, all_extra_latents, all_infos
batch = collate_rows_from_parquet_schema(rows, self.parquet_schema,
self.text_padding_length)
return batch
def __len__(self):
+13 -36
View File
@@ -411,29 +411,6 @@ 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
@@ -484,30 +461,30 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
self.text_encoding_stage = TextEncodingStage(
tokenizer=tokenizer,
text_max_length=args.text_max_length,
cfg_rate=args.training_cfg_rate)
cfg_rate=args.cfg)
def _load_raw_data(self) -> List[Dict]:
"""Load raw data from JSON files."""
all_data = []
# Read folder-annotation pairs
with open(self.data_merge_path) as f:
folder_anno_pairs = [
line.strip().split(",") for line in f if line.strip()
]
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)
# Process each folder-annotation pair
for folder, annotation_file in folder_anno_pairs:
with open(annotation_file) as f:
data_items = json.load(f)
# Update paths with folder prefix
for item in data_items:
item["path"] = opj(folder, item["path"])
# Update paths with folder prefix
for item in data_items:
item["path"] = opj(folder, item["path"])
return data_items
all_data.extend(data_items)
return all_data[self.start_idx:]
def _process_metadata(self) -> List[PreprocessBatch]:
"""Process the raw metadata through all filtering stages."""
+94 -67
View File
@@ -1,9 +1,12 @@
import random
from typing import Any, Dict, List, cast
from typing import Any, Dict, List
import numpy as np
import torch
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
"""
@@ -21,7 +24,7 @@ def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
return t[:padding_length], torch.ones(padding_length)
def get_torch_tensors_from_row_dict(row_dict, keys, cfg_rate) -> Dict[str, Any]:
def get_torch_tensors_from_row_dict(row_dict, keys) -> Dict[str, Any]:
"""
Get the latents and prompts from a row dictionary.
"""
@@ -39,45 +42,70 @@ def get_torch_tensors_from_row_dict(row_dict, keys, cfg_rate) -> Dict[str, Any]:
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"]
try:
shape = row_dict[f"{key}_shape"]
bytes = row_dict[f"{key}_bytes"]
except KeyError:
continue
# TODO (peiyuan): read precision
if key == 'text_embedding' and random.random() < cfg_rate:
data = np.zeros((512, 4096), dtype=np.float32)
if len(bytes) == 0:
return_dict[key] = torch.zeros(0, dtype=torch.bfloat16)
else:
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
data = torch.from_numpy(data)
if len(data.shape) == 3:
B, L, D = data.shape
assert B == 1, "Batch size must be 1"
data = data.squeeze(0)
return_dict[key] = data
data = torch.from_numpy(data)
if len(data.shape) == 3:
B, L, D = data.shape
assert B == 1, "Batch size must be 1"
data = data.squeeze(0)
return_dict[key] = data
return return_dict
def collate_latents_embs_masks(
batch_to_process,
text_padding_length,
keys,
cfg_rate=0.0
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, List[str]]:
batch_to_process, text_padding_length, keys
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, List[str], Dict[str, Any],
List[Dict[str, Any]]]:
# Initialize tensors to hold padded embeddings and masks
all_latents = []
all_embs = []
all_masks = []
all_clip_features = []
all_first_frame_latents = []
all_pil_images = []
all_infos = []
caption_text = []
# Process each row individually
for i, row in enumerate(batch_to_process):
# Get info from row
info_keys = [
"caption", "file_name", "media_type", "width", "height",
"num_frames", "duration_sec", "fps"
]
info = {}
for key in info_keys:
if key in row:
info[key] = row[key]
else:
info[key] = ""
info["prompt"] = info["caption"]
# Get tensors from row
data = get_torch_tensors_from_row_dict(row, keys, cfg_rate)
data = get_torch_tensors_from_row_dict(row, keys)
latents, emb = data["vae_latent"], data["text_embedding"]
clip_feature = data.get("clip_feature", None)
first_frame_latent = data.get("first_frame_latent", None)
pil_image = data.get("pil_image", None)
padded_emb, mask = pad(emb, text_padding_length)
# Store in batch tensors
all_latents.append(latents)
all_embs.append(padded_emb)
all_masks.append(mask)
all_clip_features.append(clip_feature)
all_first_frame_latents.append(first_frame_latent)
all_pil_images.append(pil_image)
all_infos.append(info)
# TODO(py): remove this once we fix preprocess
try:
caption_text.append(row["prompt"])
@@ -88,14 +116,17 @@ def collate_latents_embs_masks(
all_latents = torch.stack(all_latents)
all_embs = torch.stack(all_embs)
all_masks = torch.stack(all_masks)
all_extra_latents = {
"clip_feature": torch.stack(all_clip_features),
"first_frame_latent": torch.stack(all_first_frame_latents),
"pil_image": all_pil_images,
}
return all_latents, all_embs, all_masks, caption_text
return all_latents, all_embs, all_masks, caption_text, all_extra_latents, all_infos
def collate_rows_from_parquet_schema(rows,
parquet_schema,
text_padding_length,
cfg_rate=0.0) -> Dict[str, Any]:
def collate_rows_from_parquet_schema(rows, parquet_schema,
text_padding_length) -> Dict[str, Any]:
"""
Collate rows from parquet files based on the provided schema.
Dynamically processes tensor fields based on schema and returns batched data.
@@ -108,10 +139,10 @@ def collate_rows_from_parquet_schema(rows,
Dict containing batched tensors and metadata
"""
if not rows:
return cast(Dict[str, Any], {})
return {}
# Initialize containers for different data types
batch_data: Dict[str, Any] = {}
batch_data = {}
# Get tensor and metadata field names from schema (fields ending with '_bytes')
tensor_fields = []
@@ -128,7 +159,7 @@ def collate_rows_from_parquet_schema(rows,
# Only add actual metadata fields, not the shape/dtype helper fields
metadata_fields.append(field)
# Process each tensor field
# Process each tensor field efficiently
for tensor_name in tensor_fields:
tensor_list = []
@@ -138,6 +169,9 @@ def collate_rows_from_parquet_schema(rows,
bytes_key = f"{tensor_name}_bytes"
if shape_key in row and bytes_key in row:
# logger.info("row: %s", row)
# logger.info("shape_key: %s", shape_key)
# logger.info("bytes_key: %s", bytes_key)
shape = row[shape_key]
bytes_data = row[bytes_key]
@@ -145,12 +179,11 @@ def collate_rows_from_parquet_schema(rows,
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()
# logger.info("len(bytes_data): %s", len(bytes_data))
# logger.info("shape: %s", shape)
data = np.frombuffer(
bytes_data, dtype=np.float32).reshape(shape).copy()
tensor = torch.from_numpy(data)
# if len(data.shape) == 3:
# B, L, D = tensor.shape
@@ -163,44 +196,38 @@ def collate_rows_from_parquet_schema(rows,
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 = []
if tensor_list:
if tensor_name == 'text_embedding':
# Handle text embeddings with padding
padded_tensors = []
attention_masks = []
for tensor in tensor_list:
if tensor.numel() > 0:
padded_tensor, mask = pad(tensor, text_padding_length)
padded_tensors.append(padded_tensor)
attention_masks.append(mask)
else:
# Handle empty embeddings - assume default embedding dimension
padded_tensors.append(
torch.zeros(text_padding_length,
768,
dtype=torch.bfloat16))
attention_masks.append(torch.zeros(text_padding_length))
for tensor in tensor_list:
if tensor.numel() > 0:
padded_tensor, mask = pad(tensor, text_padding_length)
padded_tensors.append(padded_tensor)
attention_masks.append(mask)
else:
# Handle empty embeddings - assume default embedding dimension
padded_tensors.append(
torch.zeros(text_padding_length,
768,
dtype=torch.bfloat16))
attention_masks.append(torch.zeros(text_padding_length))
batch_data[tensor_name] = torch.stack(padded_tensors)
batch_data['text_attention_mask'] = torch.stack(attention_masks)
else:
# Stack 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
batch_data[tensor_name] = torch.stack(padded_tensors)
batch_data['text_attention_mask'] = torch.stack(attention_masks)
else:
# Stack other tensors directly, handling None values
valid_tensors = [
t for t in tensor_list if t is not None and t.numel() > 0
]
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
if valid_tensors:
batch_data[tensor_name] = torch.stack(valid_tensors)
elif tensor_list: # All tensors are empty but exist
batch_data[tensor_name] = torch.stack(tensor_list)
# Process metadata fields into info_list
# Process metadata fields efficiently into info_list
info_list = []
for row in rows:
info = {}
@@ -19,6 +19,7 @@ class ValidationDataset(torch.utils.data.IterableDataset):
self.filename = pathlib.Path(filename)
# get directory of filename
# TODO(will)
self.dir = os.path.abspath(self.filename.parent)
if not self.filename.exists():
+2 -2
View File
@@ -384,7 +384,7 @@ class TrainingArgs(FastVideoArgs):
# diffusion setting
ema_decay: float = 0.0
ema_start_step: int = 0
training_cfg_rate: float = 0.0
cfg: float = 0.0
precondition_outputs: bool = False
# validation & logs
@@ -528,7 +528,7 @@ class TrainingArgs(FastVideoArgs):
type=int,
default=0,
help="Step to start EMA")
parser.add_argument("--training-cfg-rate",
parser.add_argument("--cfg",
type=float,
help="Classifier-free guidance scale")
parser.add_argument(
-6
View File
@@ -38,12 +38,6 @@ class RMSNorm(CustomOp):
if self.has_weight:
self.weight = nn.Parameter(self.weight)
# if we do fully_shard(model.layer_norm), and we call layer_form.forward_native(input) instead of layer_norm(input),
# we need to call model.layer_norm.register_fsdp_forward_method(model, "forward_native") to make sure fsdp2 hooks are triggered
# for mixed precision and cpu offloading
# the even better way might be fully_shard(model.layer_norm, mp_policy=, cpu_offloading=), and call model.layer_norm(input). everything should work out of the box
# because fsdp2 hooks will be triggered with model.layer_norm.__call__
def forward_native(
self,
x: torch.Tensor,
-2
View File
@@ -14,7 +14,6 @@ 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
@@ -79,7 +78,6 @@ 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
-2
View File
@@ -442,8 +442,6 @@ 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]):
-2
View File
@@ -460,8 +460,6 @@ class StepVideoModel(BaseDiT):
# lambda n, m: "pos_embed" in n # If needed for the patch embedding.
]
_param_names_mapping = StepVideoConfig()._param_names_mapping
_reverse_param_names_mapping = StepVideoConfig(
)._reverse_param_names_mapping
_lora_param_names_mapping = StepVideoConfig()._lora_param_names_mapping
_supported_attention_backends = StepVideoConfig(
)._supported_attention_backends
+3 -1
View File
@@ -25,9 +25,12 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
PatchEmbed, TimestepEmbedder)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.dits.base import CachableDiT
from fastvideo.v1.platforms import AttentionBackendEnum
logger = init_logger(__name__)
class WanImageEmbedding(torch.nn.Module):
@@ -518,7 +521,6 @@ class WanTransformer3DModel(CachableDiT):
_supported_attention_backends = WanVideoConfig(
)._supported_attention_backends
_param_names_mapping = WanVideoConfig()._param_names_mapping
_reverse_param_names_mapping = WanVideoConfig()._reverse_param_names_mapping
_lora_param_names_mapping = WanVideoConfig()._lora_param_names_mapping
def __init__(self, config: WanVideoConfig, hf_config: dict[str,
+1 -6
View File
@@ -222,14 +222,10 @@ def load_model_from_full_model_state_dict(
used_keys = set()
sharded_sd = {}
to_merge_params: DefaultDict[str, Dict[Any, Any]] = defaultdict(dict)
reverse_param_names_mapping = {}
assert param_names_mapping is not None
for source_param_name, full_tensor in full_sd_iterator:
assert param_names_mapping is not None
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
source_param_name)
reverse_param_names_mapping[target_param_name] = (source_param_name,
merge_index,
num_params_to_merge)
used_keys.add(target_param_name)
if merge_index is not None:
to_merge_params[target_param_name][merge_index] = full_tensor
@@ -264,7 +260,6 @@ def load_model_from_full_model_state_dict(
sharded_tensor = sharded_tensor.cpu()
sharded_sd[target_param_name] = nn.Parameter(sharded_tensor)
model._reverse_param_names_mapping = reverse_param_names_mapping
unused_keys = set(meta_sd.keys()) - used_keys
if unused_keys:
logger.warning("Found new parameters in meta state dict: %s",
@@ -152,6 +152,7 @@ class TrainingBatch:
encoder_hidden_states: Optional[torch.Tensor] = None
encoder_attention_mask: Optional[torch.Tensor] = None
# i2v
# extra_latents: Optional[Dict[str, Any]] = None
preprocessed_image: Optional[torch.Tensor] = None
image_embeds: Optional[torch.Tensor] = None
image_latents: Optional[torch.Tensor] = None
@@ -104,7 +104,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
if strict:
raise ValueError(
f"Failed to convert tensor {tensor_name} to bytes: {e}"
) from e
)
record[field] = b'' # Empty bytes for missing data
else:
record[field] = b'' # Empty bytes for missing data
@@ -139,8 +139,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
except (ValueError, TypeError) as e:
if strict:
raise ValueError(
f"Failed to convert field {field} to int: {e}"
) from e
f"Failed to convert field {field} to int: {e}")
record[field] = 0
else:
record[field] = 0
@@ -158,7 +157,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
if strict:
raise ValueError(
f"Failed to convert field {field} to float: {e}"
) from e
)
record[field] = 0.0
else:
record[field] = 0.0
@@ -212,8 +211,8 @@ class BasePreprocessPipeline(ComposedPipelineBase):
# Log unfilled fields as warning if not in strict mode
if unfilled_fields:
logger.warning(
"Some fields were not filled and got default values: %s",
unfilled_fields)
f"Some fields were not filled and got default values: {unfilled_fields}"
)
return record
@@ -222,6 +221,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
# text_attention_mask: np.ndarray,
valid_data: Dict[str, Any],
idx: int,
extra_features: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
@@ -380,6 +380,8 @@ class BasePreprocessPipeline(ComposedPipelineBase):
# Convert tensors to numpy arrays
vae_latent = latent.cpu().numpy()
text_embedding = prompt_embeds[idx].cpu().numpy()
# text_attention_mask = prompt_attention_mask[idx].cpu().numpy(
# ).astype(np.uint8)
# Get extra features for this sample if needed
sample_extra_features = {}
@@ -396,6 +398,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
# text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=idx,
extra_features=sample_extra_features)
@@ -540,13 +543,14 @@ class BasePreprocessPipeline(ComposedPipelineBase):
valid_data["text"] = [prompt]
# Create record for Parquet dataset
record = self.create_record(video_name=file_name,
vae_latent=np.array([],
dtype=np.float32),
text_embedding=text_embedding,
valid_data=valid_data,
idx=0,
extra_features=sample_extra_features)
record = self.create_record(
video_name=file_name,
vae_latent=np.array([], dtype=np.float32),
text_embedding=text_embedding,
# text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=0,
extra_features=sample_extra_features)
batch_data.append(record)
logger.info("Saved validation sample: %s", file_name)
@@ -58,6 +58,7 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
result_batch = self.image_encoding_stage(batch, fastvideo_args)
clip_features = result_batch.image_embeds[0]
# image = self.pil_to_tensor(image)
image = self.preprocess(
image,
vae_scale_factor=self.get_module("vae").spatial_compression_ratio,
@@ -83,6 +84,7 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
# TODO(will): move these to cpu at some point
self.get_module("image_encoder").to(get_torch_device())
# self.get_module("image_processor").to(get_torch_device())
self.get_module("vae").to(get_torch_device())
features = {}
@@ -185,16 +187,19 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
valid_data: Dict[str, Any],
# text_attention_mask: np.ndarray,
valid_data: Optional[Dict[str, Any]],
idx: int,
extra_features: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
"""Create a record for the Parquet dataset with CLIP features."""
record = super().create_record(video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
record = super().create_record(
video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
# text_attention_mask=text_attention_mask,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
if extra_features and "clip_feature" in extra_features:
clip_feature = extra_features["clip_feature"]
@@ -59,6 +59,12 @@ if __name__ == "__main__":
default=2,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument(
"--preprocess_text_batch_size",
type=int,
default=8,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument("--samples_per_file", type=int, default=64)
parser.add_argument("--flush_frequency",
type=int,
@@ -84,7 +90,7 @@ if __name__ == "__main__":
type=str,
default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument("--training_cfg_rate", type=float, default=0.0)
parser.add_argument("--cfg", type=float, default=0.0)
parser.add_argument(
"--output_dir",
type=str,
@@ -69,7 +69,6 @@ class EncodingStage(PipelineStage):
image = image.unsqueeze(2)
else:
# assumes image is loaded from parquet file and used for validation
image = image.transpose(1, 2)
logger.info("image: %s", image.shape)
video_condition = torch.cat([
@@ -168,5 +168,5 @@ def test_clip_encoder():
f"Pooler outputs differ significantly: mean diff = {mean_diff_pooler.item()}"
assert max_diff_hidden < 1e-1, \
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
assert max_diff_pooler < 2e-2, \
assert max_diff_pooler < 1e-2, \
f"Pooler outputs differ significantly: max diff = {max_diff_pooler.item()}"
@@ -5,7 +5,7 @@ from pathlib import Path
import pytest
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "2"
NUM_GPUS_PER_NODE = "1"
# Set environment variables
os.environ["FASTVIDEO_ATTENTION_CONFIG"] = "assets/mask_strategy_wan.json"
@@ -17,9 +17,9 @@ def test_inference():
cmd = [
"fastvideo", "generate",
"--model-path", "Wan-AI/Wan2.1-T2V-14B-Diffusers",
"--sp-size", "2",
"--tp-size", "2",
"--num-gpus", "2",
"--sp-size", "1",
"--tp-size", "1",
"--num-gpus", "1",
"--height", "768",
"--width", "1280",
"--num-frames", "69",
+61 -49
View File
@@ -4,43 +4,31 @@ app = modal.App()
import os
image_version = os.getenv("IMAGE_VERSION")
image_version = os.getenv("IMAGE_VERSION", "latest")
image_tag = f"ghcr.io/hao-ai-lab/fastvideo/fastvideo-dev:{image_version}"
print(f"Using image: {image_tag}")
image = (
modal.Image.from_registry(image_tag, add_python="3.12")
.run_commands("rm -rf /FastVideo")
.apt_install("cmake", "pkg-config", "build-essential", "curl", "libssl-dev")
.run_commands("curl --proto '=https' --tlsv1.2 -sSf https://sh.rustup.rs | sh -s -- -y --default-toolchain stable")
.run_commands("echo 'source ~/.cargo/env' >> ~/.bashrc")
.env({
"PATH": "/root/.cargo/bin:$PATH",
"BUILDKITE_REPO": os.environ.get("BUILDKITE_REPO", ""),
"BUILDKITE_COMMIT": os.environ.get("BUILDKITE_COMMIT", ""),
})
.env({"PATH": "/root/.cargo/bin:$PATH"})
.run_commands("/bin/bash -c 'source $HOME/.local/bin/env && source /opt/venv/bin/activate && cd /FastVideo && uv pip install -e .[test]'")
)
def run_test(pytest_command: str):
"""Helper function to run a test suite with custom pytest command"""
@app.function(gpu="L40S:1", image=image, timeout=1800)
def run_encoder_tests():
"""Run encoder tests on L40S GPU"""
import subprocess
import sys
import os
git_repo = os.environ.get("BUILDKITE_REPO")
git_commit = os.environ.get("BUILDKITE_COMMIT")
os.chdir("/FastVideo")
print(f"Cloning repository: {git_repo}")
print(f"Checking out commit: {git_commit}")
command = f"""
source $HOME/.local/bin/env &&
source /opt/venv/bin/activate &&
git clone {git_repo} /FastVideo &&
cd /FastVideo &&
git checkout {git_commit} &&
uv pip install -e .[test] &&
{pytest_command}
command = """
source /opt/venv/bin/activate &&
pytest ./fastvideo/v1/tests/encoders -s
"""
result = subprocess.run([
@@ -49,38 +37,62 @@ def run_test(pytest_command: str):
sys.exit(result.returncode)
@app.function(gpu="L40S:1", image=image, timeout=1800)
def run_encoder_tests():
run_test("pytest ./fastvideo/v1/tests/encoders -vs")
@app.function(gpu="L40S:1", image=image, timeout=1800)
def run_vae_tests():
run_test("pytest ./fastvideo/v1/tests/vaes -vs")
"""Run VAE tests on L40S GPU"""
import subprocess
import sys
import os
os.chdir("/FastVideo")
command = """
source /opt/venv/bin/activate &&
pytest ./fastvideo/v1/tests/vaes -s
"""
result = subprocess.run([
"/bin/bash", "-c", command
], stdout=sys.stdout, stderr=sys.stderr, check=False)
sys.exit(result.returncode)
@app.function(gpu="L40S:1", image=image, timeout=1800)
def run_transformer_tests():
run_test("pytest ./fastvideo/v1/tests/transformers -vs")
"""Run transformer tests on L40S GPU"""
import subprocess
import sys
import os
os.chdir("/FastVideo")
command = """
source /opt/venv/bin/activate &&
pytest ./fastvideo/v1/tests/transformers -s
"""
result = subprocess.run([
"/bin/bash", "-c", command
], stdout=sys.stdout, stderr=sys.stderr, check=False)
sys.exit(result.returncode)
@app.function(gpu="L40S:2", image=image, timeout=3600)
def run_ssim_tests():
run_test("pytest ./fastvideo/v1/tests/ssim -vs")
@app.function(gpu="L40S:4", image=image, timeout=1800, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
def run_training_tests():
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/v1/tests/training/Vanilla -srP")
@app.function(gpu="H100:2", image=image, timeout=1800, secrets=[modal.Secret.from_dict({"WANDB_API_KEY": os.environ.get("WANDB_API_KEY", "")})])
def run_training_tests_VSA():
run_test("wandb login $WANDB_API_KEY && pytest ./fastvideo/v1/tests/training/VSA -srP")
@app.function(gpu="H100:2", image=image, timeout=1800)
def run_inference_tests_STA():
run_test("pytest ./fastvideo/v1/tests/inference/STA -srP")
@app.function(gpu="H100:1", image=image, timeout=1800)
def run_precision_tests_STA():
run_test("python csrc/attn/tests/test_sta.py")
@app.function(gpu="H100:1", image=image, timeout=1800)
def run_precision_tests_VSA():
run_test("python csrc/attn/tests/test_block_sparse.py")
"""Run SSIM tests on 2x L40S GPUs"""
import subprocess
import sys
import os
os.chdir("/FastVideo")
command = """
source /opt/venv/bin/activate &&
pytest ./fastvideo/v1/tests/ssim -vs
"""
result = subprocess.run([
"/bin/bash", "-c", command
], stdout=sys.stdout, stderr=sys.stderr, check=False)
sys.exit(result.returncode)
@@ -122,7 +122,7 @@ def run_training():
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--training_cfg_rate", "0.1",
"--cfg", "0.0",
"--output_dir", LOCAL_OUTPUT_DIR,
"--tracker_project_name", "wan_i2v_finetune_overfit_ci",
"--num_height", "480",
@@ -31,9 +31,12 @@ LOCAL_OUTPUT_DIR = Path(os.path.join(DATA_DIR, "outputs"))
def download_data():
# create the data dir if it doesn't exist
data_dir = Path(DATA_DIR)
if data_dir.exists():
print(f"Removing existing data directory at {data_dir}")
shutil.rmtree(data_dir)
print(f"Creating data directory at {data_dir}")
os.makedirs(data_dir, exist_ok=True)
os.makedirs(data_dir)
print(f"Downloading raw dataset to {LOCAL_RAW_DATA_DIR}...")
try:
@@ -119,7 +122,7 @@ def run_training():
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--training_cfg_rate", "0.0",
"--cfg", "0.0",
"--output_dir", LOCAL_OUTPUT_DIR,
"--tracker_project_name", "wan_finetune_overfit_ci",
"--num_height", "480",
+6 -7
View File
@@ -1,14 +1,13 @@
The reference videos in the `*_reference_videos` directory are used as part of an e2e test to ensure consistency in video generation quality across code changes. `test_inference_similarity.py` compares newly generated videos against these references using Structural Similarity Index (SSIM) metrics to detect any regressions in visual quality across code changes.
The reference videos in the `reference_videos` directory are used as part of an e2e test to ensure consistency in video generation quality across code changes. `test_inference_similarity.py` compares newly generated videos against these references using Structural Similarity Index (SSIM) metrics to detect any regressions in visual quality across code changes.
`A40_reference_videos` are generated on A40s and so on.
run `bash update_reference_videos.sh` from inside the `fastvideo/v1/tests/ssim/` directory after running `test_inference_similarity.py` to update reference videos. Note: make sure to update the path to the corresponding device.
all reference videos are were generated on commit `4aeabbc629e0edf91477e80e795e7bb1823c71cb`
`reference_videos/FastHunyuan-diffusers/FLASH_ATTN/` videos were generated on commit `66107fd5b8469fed25972feb632cd48887dac451`.
`reference_videos/FastHunyuan-diffusers/TORCH_SDPA/` videos were generated on commit `4ea008b8a16d7f5678a44b187ebdd7d9d0416ff1`.
`reference_videos/Wan2.1-T2V-1.3B-Diffusers` videos were generated on commit `d085770a70988c7b26632a0c3123c24a57f7ca77`.
`reference_videos/Wan2.1-I2V-14B-480P-Diffusers` videos were generated on commit `d085770a70988c7b26632a0c3123c24a57f7ca77`.
## Generation Details
2 x NVIDIA L40S GPUs
2 x NVIDIA A40 GPUs
## Generation Parameters
@@ -2,7 +2,6 @@
import json
import os
import torch
import pytest
from fastvideo import VideoGenerator
@@ -12,14 +11,6 @@ from fastvideo.v1.worker.multiproc_executor import MultiprocExecutor
logger = init_logger(__name__)
device_name = torch.cuda.get_device_name()
device_reference_folder_suffix = '_reference_videos'
if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
# Base parameters from the shell script
HUNYUAN_PARAMS = {
"num_gpus": 2,
@@ -197,8 +188,8 @@ def test_i2v_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
assert os.path.exists(
output_dir), f"Output video was not generated at {output_dir}"
reference_folder = os.path.join(script_dir, device_reference_folder, model_id, ATTENTION_BACKEND)
reference_folder = os.path.join(script_dir, 'reference_videos', model_id, ATTENTION_BACKEND)
if not os.path.exists(reference_folder):
logger.error("Reference folder missing")
raise FileNotFoundError(
@@ -297,8 +288,8 @@ def test_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
assert os.path.exists(
output_dir), f"Output video was not generated at {output_dir}"
reference_folder = os.path.join(script_dir, device_reference_folder, model_id, ATTENTION_BACKEND)
reference_folder = os.path.join(script_dir, 'reference_videos', model_id, ATTENTION_BACKEND)
if not os.path.exists(reference_folder):
logger.error("Reference folder missing")
raise FileNotFoundError(
@@ -1,63 +0,0 @@
#!/bin/bash
# Script to update reference videos using videos from generated_videos directory
# Both directories should exist in the same directory as this script
set -e # Exit on any error
# Define directory paths
GENERATED_DIR="generated_videos"
REFERENCE_DIR="set_me_to_correct_path"
# Colors for output
RED='\033[0;31m'
GREEN='\033[0;32m'
YELLOW='\033[1;33m'
NC='\033[0m' # No Color
echo -e "${YELLOW}Starting reference video update...${NC}"
# Check if generated_videos directory exists
if [ ! -d "$GENERATED_DIR" ]; then
echo -e "${RED}Error: $GENERATED_DIR directory not found!${NC}"
exit 1
fi
# Check if reference_videos directory exists
if [ ! -d "$REFERENCE_DIR" ]; then
echo -e "${RED}Error: $REFERENCE_DIR directory not found!${NC}"
exit 1
fi
# Function to copy videos recursively
copy_videos() {
local src_dir="$1"
local dst_dir="$2"
# Find all video files in the source directory
find "$src_dir" -type f \( -name "*.mp4" -o -name "*.avi" -o -name "*.mov" -o -name "*.mkv" -o -name "*.webm" -o -name "*.flv" \) | while read -r video_file; do
# Get relative path from source directory
relative_path="${video_file#$src_dir/}"
# Construct destination path
dst_file="$dst_dir/$relative_path"
# Create destination directory if it doesn't exist
dst_file_dir=$(dirname "$dst_file")
mkdir -p "$dst_file_dir"
# Copy the video file
echo -e "${GREEN}Copying: $relative_path${NC}"
cp "$video_file" "$dst_file"
done
}
# Perform the copy operation
echo -e "${YELLOW}Copying videos from $GENERATED_DIR to $REFERENCE_DIR...${NC}"
copy_videos "$GENERATED_DIR" "$REFERENCE_DIR"
echo -e "${GREEN}Reference videos updated successfully!${NC}"
# Show summary
video_count=$(find "$GENERATED_DIR" -type f \( -name "*.mp4" -o -name "*.avi" -o -name "*.mov" -o -name "*.mkv" -o -name "*.webm" -o -name "*.flv" \) | wc -l)
echo -e "${YELLOW}Total videos processed: $video_count${NC}"
@@ -1 +1 @@
{"step_time":0.6983645600266755,"_wandb":{"runtime":107},"grad_norm":0.50390625,"avg_step_time":1.002151239803061,"_step":5,"validation_videos_50_steps":{"captions":false,"_type":"videos","count":8,"videos":[{"size":159131,"path":"media/videos/validation_videos_50_steps_0_dc447599dbe48350e9c9.mp4","_type":"video-file","sha256":"dc447599dbe48350e9c920f4971e1786bde580dd20d52b4aa147ae8d3dc564d6"},{"_type":"video-file","sha256":"4e283876ddfbf5a2cb6f8aca07a39f832b5f806fbefde8854bc73ed904ff20ee","size":160315,"path":"media/videos/validation_videos_50_steps_0_4e283876ddfbf5a2cb6f.mp4"},{"size":135225,"path":"media/videos/validation_videos_50_steps_0_78185c41e1935306e93c.mp4","_type":"video-file","sha256":"78185c41e1935306e93c2d416ee40b31d038abffb029cb5bfb11c2a634eb2fcf"},{"_type":"video-file","sha256":"27e9819d002d3f63c8918bbdc5bf2857b5effe0caf5d5b8374b9c590fc6432eb","size":197873,"path":"media/videos/validation_videos_50_steps_0_27e9819d002d3f63c891.mp4"},{"_type":"video-file","sha256":"46fe548e86144ca60a9396fcd15e8788d3a05c066c6819ccdd6a9041feaaec8f","size":170601,"path":"media/videos/validation_videos_50_steps_0_46fe548e86144ca60a93.mp4"},{"sha256":"91ec338774bec870b9c5c81be4330a8f0f2535124cb372e881c3703e6b65ed77","size":164462,"path":"media/videos/validation_videos_50_steps_0_91ec338774bec870b9c5.mp4","_type":"video-file"},{"_type":"video-file","sha256":"ee4e811080a619215fd7541b39203dfb69c2db4cdadbfdd7a234a2cff684f6f3","size":139435,"path":"media/videos/validation_videos_50_steps_0_ee4e811080a619215fd7.mp4"},{"sha256":"22e31e048ba5e5b9d6587d7306904d6edc6c618b8f77c3ceb05ba2d602309274","size":147072,"path":"media/videos/validation_videos_50_steps_0_22e31e048ba5e5b9d658.mp4","_type":"video-file"}]},"_timestamp":1.75118195270901e+09,"vsa_sparsity":0.05,"learning_rate":1e-05,"train_loss":0.19960195198655128,"_runtime":107.325113071}
{"grad_norm":0.478515625,"_runtime":95.727033597,"_wandb":{"runtime":95},"_step":5,"validation_videos_50_steps":{"videos":[{"_type":"video-file","sha256":"42a1c311521a9d460db788713be1cbf2db767494e02619b43be5bf3eed8381d8","size":158632,"path":"media/videos/validation_videos_50_steps_0_42a1c311521a9d460db7.mp4"},{"path":"media/videos/validation_videos_50_steps_0_818505095b4b5e8b7f51.mp4","_type":"video-file","sha256":"818505095b4b5e8b7f511012d45f04d151ce3344bc058fc0f3225a414a851e4a","size":147825},{"sha256":"fc334ba9ed5e66c8527ee3b408e3be2d76167fef03588bf2840f4a0792f2fe34","size":136933,"path":"media/videos/validation_videos_50_steps_0_fc334ba9ed5e66c8527e.mp4","_type":"video-file"},{"size":201797,"path":"media/videos/validation_videos_50_steps_0_ccd98f6f907635d266a7.mp4","_type":"video-file","sha256":"ccd98f6f907635d266a74783688e7ecf1dac752d79d72d69eab9ef0e3f7413eb"},{"_type":"video-file","sha256":"ca79f40a0aed38f676f12779b349ce40e9e3fb7f36c578f49a20854c70508fb4","size":147114,"path":"media/videos/validation_videos_50_steps_0_ca79f40a0aed38f676f1.mp4"},{"size":175104,"path":"media/videos/validation_videos_50_steps_0_32c9b33ff920c17e5881.mp4","_type":"video-file","sha256":"32c9b33ff920c17e588133d7a27aa400ff3dc529b01ed4f16ac4d6bb2afa0f00"},{"sha256":"2cf520bfb93401c914e93c87ef791c2f12a4e043b95dfdc98115c930e11dfe67","size":139655,"path":"media/videos/validation_videos_50_steps_0_2cf520bfb93401c914e9.mp4","_type":"video-file"},{"_type":"video-file","sha256":"1d73aba17ce582c7aef4af4d64079e3e9d3df205634eff453446bdaf2340b214","size":149028,"path":"media/videos/validation_videos_50_steps_0_1d73aba17ce582c7aef4.mp4"}],"captions":false,"_type":"videos","count":8},"train_loss":0.08922439813613892,"_timestamp":1.750202051751466e+09,"avg_step_time":0.7536672964692116,"step_time":0.4742048177868128,"learning_rate":1e-05,"vsa_sparsity":0.05}
@@ -15,7 +15,7 @@ wandb_name = "test_training_loss_VSA"
reference_wandb_summary_file = "fastvideo/v1/tests/training/VSA/reference_wandb_summary_VSA.json"
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "2"
NUM_GPUS_PER_NODE = "1"
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VIDEO_SPARSE_ATTN"
@@ -31,18 +31,19 @@ def run_worker():
"--model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--inference_mode", "False",
"--pretrained_model_name_or_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--cache_dir", "/home/.cache",
"--data_path", "data/mini_dataset_i2v_VSA/combined_parquet_dataset",
"--validation_preprocessed_path", "data/mini_dataset_i2v_VSA/validation_parquet_dataset",
"--train_batch_size", "1",
"--num_latent_t", "4",
"--num_gpus", "2",
"--sp_size", "2",
"--tp_size", "2",
"--num_gpus", "1",
"--sp_size", "1",
"--tp_size", "1",
"--hsdp_replicate_dim", "1",
"--hsdp_shard_dim", "2",
"--hsdp_shard_dim", "1",
"--train_sp_batch_size", "1",
"--dataloader_num_workers", "4",
"--gradient_accumulation_steps", "2",
"--gradient_accumulation_steps", "1",
"--max_train_steps", "5",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
@@ -53,7 +54,7 @@ def run_worker():
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--training_cfg_rate", "0.0",
"--cfg", "0.0",
"--output_dir", "data/wan_finetune_test_VSA",
"--tracker_project_name", "wan_finetune_ci_VSA",
"--wandb_run_name", wandb_name,
@@ -110,7 +111,7 @@ def test_distributed_training():
fields_and_thresholds = {
'avg_step_time': 1.0,
'grad_norm': 0.1,
'step_time': 1.0,
'step_time': 0.5,
'train_loss': 0.001
}
@@ -1 +0,0 @@
{"step_time":5.501357046999999,"grad_norm":0.384765625,"train_loss":0.07890288904309273,"avg_step_time":5.831571423200001}
@@ -1 +1 @@
{"_timestamp":1.7496170016478686e+09,"validation_videos_8_steps":{"_type":"videos","count":5,"videos":[{"sha256":"d81aa715df0c3ba4db8b242bc10442844453f3684a45dc55b89ecd27b2414fc7","size":429642,"path":"media/videos/validation_videos_8_steps_0_d81aa715df0c3ba4db8b.mp4","_type":"video-file"},{"sha256":"cd72a3d513eca6b41b03b80e6fa044ce7219c35e969d2ca20b9cb48c91e585c6","size":477837,"path":"media/videos/validation_videos_8_steps_0_cd72a3d513eca6b41b03.mp4","_type":"video-file"},{"_type":"video-file","sha256":"43d47c211a69bf0be3544738e76e7d8fa58bb108c00cb21dc34eaad3c6ce7cc3","size":409419,"path":"media/videos/validation_videos_8_steps_0_43d47c211a69bf0be354.mp4"},{"_type":"video-file","sha256":"ea674ec9e200bc97563c9d87d9dc07110c3f42c0ab7277dd237ae96dd8f90a10","size":333966,"path":"media/videos/validation_videos_8_steps_0_ea674ec9e200bc97563c.mp4"},{"sha256":"d81aa715df0c3ba4db8b242bc10442844453f3684a45dc55b89ecd27b2414fc7","size":429642,"path":"media/videos/validation_videos_8_steps_0_d81aa715df0c3ba4db8b.mp4","_type":"video-file"}],"captions":false},"step_time":2.5065076276659966,"_wandb":{"runtime":53},"learning_rate":1e-06,"_step":5,"_runtime":53.172758961,"grad_norm":0.408203125,"train_loss":0.07883700542151928,"avg_step_time":2.8116052336990833}
{"_timestamp":1.7496170016478686e+09,"validation_videos_8_steps":{"_type":"videos","count":5,"videos":[{"sha256":"d81aa715df0c3ba4db8b242bc10442844453f3684a45dc55b89ecd27b2414fc7","size":429642,"path":"media/videos/validation_videos_8_steps_0_d81aa715df0c3ba4db8b.mp4","_type":"video-file"},{"sha256":"cd72a3d513eca6b41b03b80e6fa044ce7219c35e969d2ca20b9cb48c91e585c6","size":477837,"path":"media/videos/validation_videos_8_steps_0_cd72a3d513eca6b41b03.mp4","_type":"video-file"},{"_type":"video-file","sha256":"43d47c211a69bf0be3544738e76e7d8fa58bb108c00cb21dc34eaad3c6ce7cc3","size":409419,"path":"media/videos/validation_videos_8_steps_0_43d47c211a69bf0be354.mp4"},{"_type":"video-file","sha256":"ea674ec9e200bc97563c9d87d9dc07110c3f42c0ab7277dd237ae96dd8f90a10","size":333966,"path":"media/videos/validation_videos_8_steps_0_ea674ec9e200bc97563c.mp4"},{"sha256":"d81aa715df0c3ba4db8b242bc10442844453f3684a45dc55b89ecd27b2414fc7","size":429642,"path":"media/videos/validation_videos_8_steps_0_d81aa715df0c3ba4db8b.mp4","_type":"video-file"}],"captions":false},"step_time":2.5065076276659966,"_wandb":{"runtime":53},"learning_rate":1e-06,"_step":5,"_runtime":53.172758961,"grad_norm":5.65625,"train_loss":0.3915919363498688,"avg_step_time":2.8116052336990833}
@@ -18,8 +18,7 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
wandb_name = "test_training_loss"
a40_reference_wandb_summary_file = "fastvideo/v1/tests/training/Vanilla/a40_reference_wandb_summary.json"
l40s_reference_wandb_summary_file = "fastvideo/v1/tests/training/Vanilla/l40s_reference_wandb_summary.json"
reference_wandb_summary_file = "fastvideo/v1/tests/training/Vanilla/reference_wandb_summary.json"
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "4"
@@ -37,8 +36,9 @@ def run_worker():
"--model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--inference_mode", "False",
"--pretrained_model_name_or_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--data_path", "data/crush-smol_processed_t2v/combined_parquet_dataset",
"--validation_preprocessed_path", "data/crush-smol_processed_t2v/validation_parquet_dataset",
"--cache_dir", "/home/.cache",
"--data_path", "data/crush-smol_parq/combined_parquet_dataset",
"--validation_preprocessed_path", "data/crush-smol_parq/validation_parquet_dataset",
"--train_batch_size", "2",
"--num_latent_t", "4",
"--num_gpus", "4",
@@ -59,7 +59,7 @@ def run_worker():
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--training_cfg_rate", "0.0",
"--cfg", "0.0",
"--output_dir", "data/wan_finetune_test",
"--tracker_project_name", "wan_finetune_ci",
"--wandb_run_name", wandb_name,
@@ -83,12 +83,12 @@ def test_distributed_training():
"""Test the distributed training setup"""
os.environ["WANDB_MODE"] = "online"
data_dir = Path("data/crush-smol_processed_t2v")
data_dir = Path("data/crush-smol_parq")
if not data_dir.exists():
print(f"Downloading test dataset to {data_dir}...")
snapshot_download(
repo_id="wlsaidhi/crush-smol_processed_t2v",
repo_id="PY007/crush-smol",
local_dir=str(data_dir),
repo_type="dataset",
local_dir_use_symlinks=False
@@ -109,22 +109,14 @@ def test_distributed_training():
summary_file = 'wandb/latest-run/files/wandb-summary.json'
device_name = torch.cuda.get_device_name()
if "A40" in device_name:
reference_wandb_summary_file = a40_reference_wandb_summary_file
elif "L40S" in device_name:
reference_wandb_summary_file = l40s_reference_wandb_summary_file
else:
raise ValueError(f"Unknown device: {device_name}")
reference_wandb_summary = json.load(open(reference_wandb_summary_file))
wandb_summary = json.load(open(summary_file))
fields_and_thresholds = {
'avg_step_time': 6.0,
'grad_norm': 0.3,
'step_time': 6.0,
'train_loss': 0.0025
'avg_step_time': 1.0,
'grad_norm': 0.2,
'step_time': 0.5,
'train_loss': 0.001
}
failures = []
@@ -92,7 +92,7 @@ def test_hunyuanvideo_distributed():
# Move to GPU based on local rank (0 or 1 for 2 GPUs)
device = torch.device(f"cuda:0")
model = model
model = model.to(device)
batch_size = 1
seq_len = 3
@@ -62,7 +62,7 @@ def test_hunyuanvideo_distributed():
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
use_cpu_offload=True,
use_cpu_offload=False,
pipeline_config=PipelineConfig(dit_config=HunyuanVideoConfig(), dit_precision=precision_str))
args.device = torch.device(f"cuda:{LOCAL_RANK}")
@@ -34,12 +34,12 @@ def test_wan_transformer():
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
use_cpu_offload=True,
use_cpu_offload=False,
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, "", args).to(dtype=precision)
model2 = loader.load(TRANSFORMER_PATH, "", args).to(device, dtype=precision)
model1 = WanTransformer3DModel.from_pretrained(
TRANSFORMER_PATH, device=device,
+2 -14
View File
@@ -28,12 +28,8 @@ MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
VAE_PATH = os.path.join(MODEL_PATH, "vae")
CONFIG_PATH = os.path.join(VAE_PATH, "config.json")
# Latent generated on commit d71a4ebffc2034922fc379568b6a6aa722f3744c with 1 x A40
# torch 2.7.1
A40_REFERENCE_LATENT = -106.22467041015625
# Latent generated on commit 2b54068960c41d42221e8b8719a374b499855029 with 1 x L40S
L40S_REFERENCE_LATENT = -158.32318115234375
# Latent generated on commit 250f0b916cebb18a1c15c4aae1a0b480604d066a with 1 x A40
REFERENCE_LATENT = -105.51324462890625
@pytest.mark.usefixtures("distributed_setup")
@@ -70,14 +66,6 @@ def test_hunyuan_vae():
latent = model.encode(input_tensor).mean.double().sum().item()
# Check if latents are similar
device_name = torch.cuda.get_device_name()
if "A40" in device_name:
REFERENCE_LATENT = A40_REFERENCE_LATENT
elif "L40S" in device_name:
REFERENCE_LATENT = L40S_REFERENCE_LATENT
else:
raise ValueError(f"Unknown device: {device_name}")
diff_encoded_latents = abs(REFERENCE_LATENT - latent)
logger.info(
f"Reference latent: {REFERENCE_LATENT}, Current latent: {latent}"
+25 -28
View File
@@ -59,7 +59,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
raise RuntimeError(
"create_pipeline_stages should not be called for training pipeline")
def set_schemas(self) -> None:
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_t2v
self.validation_dataset_schema = pyarrow_schema_t2v_validation
@@ -79,6 +79,8 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
self.seed = training_args.seed
assert self.transformer is not None
self.set_schemas()
# self.train_dataset_schema = pyarrow_schema_t2v
# self.validation_dataset_schema = pyarrow_schema_t2v_validation
self.transformer.requires_grad_(True)
self.transformer.train()
@@ -114,7 +116,6 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
training_args.train_batch_size,
parquet_schema=self.train_dataset_schema,
num_data_workers=training_args.dataloader_num_workers,
cfg_rate=training_args.training_cfg_rate,
drop_last=True,
text_padding_length=training_args.pipeline_config.
text_encoder_configs[0].arch_config.
@@ -156,7 +157,6 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
return training_batch
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert self.train_loader_iter is not None
assert self.train_dataloader is not None
@@ -169,11 +169,18 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
# Get first batch of new epoch
batch = next(self.train_loader_iter)
# latents, encoder_hidden_states, encoder_attention_mask, infos = batch
# latents, encoder_hidden_states, encoder_attention_mask, caption_text, extra_latents, infos = batch
# for key, value in batch.items():
# if isinstance(value, torch.Tensor):
# logger.info("key: %s, shape: %s", key, value.shape)
# else:
# logger.info("key: %s, value: %s", key, value)
# print("--------------------------------")
# logger.info("batch: %s", batch)
latents = batch['vae_latent']
latents = latents[:, :, :self.training_args.num_latent_t]
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
# extra_latents = batch['extra_latents']
infos = batch['info_list']
training_batch.latents = latents.to(get_torch_device(),
@@ -182,6 +189,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
get_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_torch_device(), dtype=torch.bfloat16)
# training_batch.extra_latents = extra_latents
training_batch.infos = infos
return training_batch
@@ -245,8 +253,9 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
current_vsa_sparsity = training_batch.current_vsa_sparsity
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
dit_seq_shape = [
latents.shape[2] * self.sp_world_size // patch_size[0],
latents.shape[2] // patch_size[0],
latents.shape[3] // patch_size[1],
latents.shape[4] // patch_size[2]
]
@@ -286,10 +295,8 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.transformer is not None
assert self.training_args is not None
assert training_batch.noisy_model_input is not None
assert training_batch.latents is not None
assert training_batch.noise is not None
assert training_batch.sigmas is not None
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
assert training_batch.attn_metadata is not None
@@ -313,8 +320,6 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
target = training_batch.latents if self.training_args.precondition_outputs else training_batch.noise - training_batch.latents
# make sure no implicit broadcasting happens
assert model_pred.shape == target.shape, f"model_pred.shape: {model_pred.shape}, target.shape: {target.shape}"
loss = (torch.mean((model_pred.float() - target.float())**2) /
self.training_args.gradient_accumulation_steps)
@@ -358,24 +363,14 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
for _ in range(self.training_args.gradient_accumulation_steps):
training_batch = self._get_next_batch(training_batch)
# Normalize DIT input
training_batch = self._normalize_dit_input(training_batch)
# Create noisy model input
training_batch = self._prepare_dit_inputs(training_batch)
# Shard latents across sp groups
training_batch.latents = shard_latents_across_sp(
training_batch.latents,
num_latent_t=self.training_args.num_latent_t)
# shard noisy_model_input to match
training_batch.noisy_model_input = shard_latents_across_sp(
training_batch.noisy_model_input,
num_latent_t=self.training_args.num_latent_t)
# shard noise to match latents
training_batch.noise = shard_latents_across_sp(
training_batch.noise,
num_latent_t=self.training_args.num_latent_t)
# Normalize DIT input
training_batch = self._normalize_dit_input(training_batch)
training_batch = self._prepare_dit_inputs(training_batch)
training_batch = self._build_attention_metadata(training_batch)
training_batch = self._build_input_kwargs(training_batch)
training_batch = self._transformer_forward_and_compute_loss(
@@ -546,6 +541,8 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
negative_prompt_attention_mask: torch.Tensor | None
) -> ForwardBatch:
# logger.info("validation_batch: %s", validation_batch)
# latents, embeddings, masks, caption_text, extra_latents, infos = validation_batch
prompt = validation_batch['info_list'][0]['prompt']
prompt_embeds = validation_batch['text_embedding']
prompt_attention_mask = validation_batch['text_attention_mask']
@@ -602,6 +599,7 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
# Set deterministic seed for validation
set_random_seed(self.seed)
logger.info("Using validation seed: %s", self.seed)
# Prepare validation prompts
@@ -612,13 +610,13 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
batch_size=1,
parquet_schema=self.validation_dataset_schema,
num_data_workers=0,
cfg_rate=0.0,
drop_last=False,
drop_first_row=sampling_param.negative_prompt is not None)
drop_first_row=sampling_param.negative_prompt is not None,
cfg_rate=training_args.cfg)
if sampling_param.negative_prompt:
negative_prompt_embeds, negative_prompt_attention_mask, negative_prompt = validation_dataset.get_validation_negative_prompt(
)
logger.info("Using negative_prompt: %s", negative_prompt)
logger.info("negative_prompt: %s", negative_prompt)
transformer.eval()
@@ -631,14 +629,13 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
step_videos: List[np.ndarray] = []
step_captions: List[str | None] = []
# for _, embeddings, masks, caption_text, extra_latents, infos in validation_dataloader:
for validation_batch in validation_dataloader:
batch = self._prepare_validation_inputs(
sampling_param, training_args, validation_batch,
num_inference_steps, negative_prompt_embeds,
negative_prompt_attention_mask)
step_captions.extend([None]) # TODO(peiyuan): add caption
# Run validation inference
with torch.no_grad(), torch.autocast("cuda",
dtype=torch.bfloat16):
+4 -75
View File
@@ -8,6 +8,7 @@ from typing import Any, Dict, List, Optional, Tuple, Union
import torch
import torch.distributed as dist
import torch.distributed.checkpoint as dcp
import torch.distributed.checkpoint.stateful
from einops import rearrange
from safetensors.torch import save_file
@@ -153,20 +154,13 @@ def save_checkpoint(transformer,
if rank == 0:
# Save model weights (consolidated)
transformer_save_dir = os.path.join(save_dir, "transformer")
os.makedirs(transformer_save_dir, exist_ok=True)
weight_path = os.path.join(transformer_save_dir,
weight_path = os.path.join(save_dir,
"diffusion_pytorch_model.safetensors")
logger.info("rank: %s, saving consolidated checkpoint to %s",
rank,
weight_path,
local_main_process_only=False)
# Convert training format to diffusers format and save
diffusers_state_dict = convert_training_to_diffusers_format(
cpu_state, transformer)
save_file(diffusers_state_dict, weight_path)
save_file(cpu_state, weight_path)
logger.info("rank: %s, consolidated checkpoint saved to %s",
rank,
weight_path,
@@ -176,7 +170,7 @@ def save_checkpoint(transformer,
config_dict = transformer.hf_config
if "dtype" in config_dict:
del config_dict["dtype"] # TODO
config_path = os.path.join(transformer_save_dir, "config.json")
config_path = os.path.join(save_dir, "config.json")
# save dict as json
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
@@ -485,68 +479,3 @@ def _has_foreach_support(tensors: List[torch.Tensor],
device: torch.device) -> bool:
return _device_has_foreach_support(device) and all(
t is None or type(t) in [torch.Tensor] for t in tensors)
def convert_training_to_diffusers_format(state_dict: Dict[str, Any],
transformer) -> Dict[str, Any]:
"""
Convert training format state dict to diffusers format using reverse_param_names_mapping.
Args:
state_dict: State dict in training format
transformer: Transformer model object with _reverse_param_names_mapping
Returns:
State dict in diffusers format
"""
new_state_dict = {}
# Get the reverse mapping from the transformer
reverse_param_names_mapping = transformer._reverse_param_names_mapping
assert reverse_param_names_mapping != {}, "reverse_param_names_mapping is empty"
# Group parameters that need to be split (merged parameters)
merge_groups: Dict[str, List[Tuple[str, int, int]]] = {}
# First pass: collect all merge groups
for training_key, (
diffusers_key, merge_index,
num_params_to_merge) in reverse_param_names_mapping.items():
if merge_index is not None:
# This is a merged parameter that needs to be split
if training_key not in merge_groups:
merge_groups[training_key] = []
merge_groups[training_key].append(
(diffusers_key, merge_index, num_params_to_merge))
# Second pass: handle merged parameters by splitting them
used_keys = set()
for training_key, splits in merge_groups.items():
if training_key in state_dict:
v = state_dict[training_key]
# Sort by merge_index to ensure correct order
splits.sort(key=lambda x: x[1])
total = splits[0][2]
split_size = v.shape[0] // total
split_tensors = torch.split(v, split_size, dim=0)
for diffusers_key, split_index, _ in splits:
new_state_dict[diffusers_key] = split_tensors[split_index]
used_keys.add(training_key)
# Third pass: handle regular parameters (direct mappings)
for training_key, v in state_dict.items():
if training_key in used_keys:
continue
if training_key in reverse_param_names_mapping:
diffusers_key, merge_index, _ = reverse_param_names_mapping[
training_key]
if merge_index is None:
# Direct mapping
new_state_dict[diffusers_key] = v
else:
# No mapping found, keep as is
new_state_dict[training_key] = v
return new_state_dict
@@ -4,12 +4,11 @@ from copy import deepcopy
from typing import Any, Dict
import torch
import torch.distributed
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset.dataloader.schema import (
pyarrow_schema_i2v, pyarrow_schema_i2v_validation)
from fastvideo.v1.distributed import get_torch_device, get_sp_group
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.scheduling_flow_unipc_multistep import (
@@ -19,8 +18,7 @@ from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
from fastvideo.v1.pipelines.wan.wan_i2v_pipeline import (
WanImageToVideoValidationPipeline)
from fastvideo.v1.training.training_pipeline import TrainingPipeline
from fastvideo.v1.training.training_utils import (shard_latents_across_sp,
clip_grad_norm_while_handling_failing_dtensor_cases)
from fastvideo.v1.training.training_utils import shard_latents_across_sp
from fastvideo.v1.utils import is_vsa_available
vsa_available = is_vsa_available()
@@ -66,7 +64,7 @@ class WanI2VTrainingPipeline(TrainingPipeline):
self.validation_pipeline = validation_pipeline
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
assert self.train_loader_iter is not None
assert self.train_dataloader is not None
batch = next(self.train_loader_iter, None) # type: ignore
@@ -78,14 +76,21 @@ class WanI2VTrainingPipeline(TrainingPipeline):
# Get first batch of new epoch
batch = next(self.train_loader_iter)
# latents, encoder_hidden_states, encoder_attention_mask, caption_text, extra_latents, infos = batch
# for key, value in batch.items():
# if isinstance(value, torch.Tensor):
# logger.info("key: %s, shape: %s", key, value.shape)
# else:
# logger.info("key: %s, value: %s", key, value)
# print("--------------------------------")
# logger.info("batch: %s", batch)
latents = batch['vae_latent']
latents = latents[:, :, :self.training_args.num_latent_t]
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
clip_features = batch['clip_feature']
image_latents = batch['first_frame_latent']
image_latents = image_latents[:, :, :self.training_args.num_latent_t]
pil_image = batch['pil_image']
# extra_latents = batch['extra_latents']
infos = batch['info_list']
training_batch.latents = latents.to(get_torch_device(),
@@ -97,32 +102,11 @@ class WanI2VTrainingPipeline(TrainingPipeline):
training_batch.preprocessed_image = pil_image.to(get_torch_device())
training_batch.image_embeds = clip_features.to(get_torch_device())
training_batch.image_latents = image_latents.to(get_torch_device())
# training_batch.extra_latents = extra_latents
training_batch.infos = infos
return training_batch
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
"""Override to properly handle I2V concatenation - call parent first, then concatenate image conditioning."""
assert self.training_args is not None
assert training_batch.latents is not None
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert self.noise_random_generator is not None
assert training_batch.image_latents is not None
# First, call parent method to prepare noise, timesteps, etc. for video latents
training_batch = super()._prepare_dit_inputs(training_batch)
assert isinstance(training_batch.image_latents, torch.Tensor)
image_latents = training_batch.image_latents.to(get_torch_device(),
dtype=torch.bfloat16)
training_batch.noisy_model_input = torch.cat(
[training_batch.noisy_model_input, image_latents], dim=1)
return training_batch
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
assert self.training_args is not None
@@ -130,15 +114,34 @@ class WanI2VTrainingPipeline(TrainingPipeline):
assert training_batch.encoder_hidden_states is not None
assert training_batch.encoder_attention_mask is not None
assert training_batch.timesteps is not None
assert training_batch.preprocessed_image is not None
assert training_batch.image_embeds is not None
assert training_batch.image_latents is not None
# assert training_batch.extra_latents is not None
# Image Embeds for conditioning
# extra_latents = training_batch.extra_latents
# if extra_latents:
# image_embeds, image_latents = extra_latents[
# "clip_feature"], extra_latents["first_frame_latent"]
# image_
# Image Embeds
image_embeds = training_batch.image_embeds
image_latents = training_batch.image_latents
preprocessed_image = training_batch.preprocessed_image
assert torch.isnan(image_embeds).sum() == 0
image_embeds = image_embeds.to(get_torch_device(), dtype=torch.bfloat16)
encoder_hidden_states_image = image_embeds
# NOTE: noisy_model_input already contains concatenated image_latents from _prepare_dit_inputs
# Image Latents
assert torch.isnan(image_latents).sum() == 0
image_latents = image_latents.to(get_torch_device(),
dtype=torch.bfloat16)
image_latents = shard_latents_across_sp(
image_latents, num_latent_t=self.training_args.num_latent_t)
training_batch.noisy_model_input = torch.cat(
[training_batch.noisy_model_input, image_latents], dim=1)
training_batch.input_kwargs = {
"hidden_states":
training_batch.noisy_model_input,
@@ -162,10 +165,13 @@ class WanI2VTrainingPipeline(TrainingPipeline):
negative_prompt_embeds: torch.Tensor | None,
negative_prompt_attention_mask: torch.Tensor | None
) -> ForwardBatch:
# latents, embeddings, masks, caption_text, extra_latents, infos = validation_batch
# latents = validation_batch['vae_latent']
embeddings = validation_batch['text_embedding']
masks = validation_batch['text_attention_mask']
clip_features = validation_batch['clip_feature']
pil_image = validation_batch['pil_image']
# extra_latents = validation_batch['extra_latents']
preprocessed_image = validation_batch['pil_image']
infos = validation_batch['info_list']
prompt = infos[0]['prompt']
@@ -173,6 +179,33 @@ class WanI2VTrainingPipeline(TrainingPipeline):
prompt_attention_mask = masks.to(get_torch_device())
clip_features = clip_features.to(get_torch_device())
if 'colorful candies' in prompt:
logger.info("colorful candies")
from fastvideo.v1.models.vision_utils import load_video
# video_path = 'validation_dataset/yYcK4nANZz4-Scene-030.mp4'
video_path = '/mnt/user_storage/fv/FastVideo/examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-030.mp4'
video = load_video(video_path)
pil_image = video[0]
preprocessed_image = None
else:
pil_image = None
# clip_features = extra_latents.get("clip_feature")
# first_frame_latent = extra_latents.get("first_frame_latent")
# pil_image = extra_latents.get("pil_image")
# if clip_features is not None and clip_features.numel() > 0:
# clip_features = clip_features.to(get_torch_device())
# if first_frame_latent is not None and first_frame_latent.numel() > 0:
# first_frame_latent = first_frame_latent.to(get_torch_device())
# if pil_image is not None and pil_image[0] is not None and pil_image[
# 0].numel() > 0:
# pil_image = pil_image[0].to(get_torch_device())
# else:
# clip_features = None
# first_frame_latent = None
# pil_image = None
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
@@ -194,7 +227,8 @@ class WanI2VTrainingPipeline(TrainingPipeline):
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
image_embeds=[clip_features],
preprocessed_image=pil_image,
preprocessed_image=preprocessed_image,
pil_image=pil_image,
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
@@ -207,36 +241,6 @@ class WanI2VTrainingPipeline(TrainingPipeline):
)
return batch
def _clip_grad_norm(self, training_batch: TrainingBatch) -> TrainingBatch:
"""Override to add gradient synchronization across SP ranks."""
assert self.training_args is not None
max_grad_norm = self.training_args.max_grad_norm
# CRITICAL FIX: Synchronize gradients across SP ranks before clipping
# Different SP ranks compute different gradients due to different noise patterns
# These gradients must be averaged across SP ranks for stable training
if self.training_args.sp_size > 1:
sp_group = get_sp_group()
for param in self.transformer.parameters():
if param.grad is not None:
# Average gradients across SP ranks
sp_group.all_reduce(param.grad, op=torch.distributed.ReduceOp.AVG)
if max_grad_norm is not None:
model_parts = [self.transformer]
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for m in model_parts for p in m.parameters()],
max_grad_norm,
foreach=None,
)
assert grad_norm is not float('nan') or grad_norm is not float(
'inf')
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
else:
grad_norm = 0.0
training_batch.grad_norm = grad_norm
return training_batch
def main(args) -> None:
logger.info("Starting training pipeline...")
+1 -2
View File
@@ -1,4 +1,3 @@
# trigger test
[build-system]
requires = ["setuptools>=61.0"]
build-backend = "setuptools.build_meta"
@@ -20,7 +19,7 @@ dependencies = [
# Machine Learning & Transformers
"transformers>=4.46.1", "tokenizers>=0.20.1", "sentencepiece==0.2.0",
"timm==1.0.11", "peft>=0.15.0", "diffusers>=0.33.1", "bitsandbytes",
"timm==1.0.11", "peft==0.13.2", "diffusers>=0.33.1", "bitsandbytes",
"torch==2.7.1", "torchvision",
# Acceleration & Optimization
+1 -1
View File
@@ -36,7 +36,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--checkpoints_total_limit 3\
--allow_tf32\
--ema_start_step 0\
--training_cfg_rate 0.0\
--cfg 0.0\
--output_dir="$DATA_DIR/outputs/wan_finetune"\
--tracker_project_name wan_finetune \
--num_height 480 \
+1 -1
View File
@@ -42,7 +42,7 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
--checkpoints_total_limit 3 \
--allow_tf32 \
--ema_start_step 0 \
--training_cfg_rate 0.0 \
--cfg 0.0 \
--output_dir "$DATA_DIR/outputs/wan_finetune" \
--tracker_project_name VSA_finetune \
--num_height 448 \
-27
View File
@@ -1,27 +0,0 @@
#!/bin/bash
num_gpus=1
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
# change model path to local dir if you want to inference using your checkpoint
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# Note that the tp_size and sp_size should be the same and equal to the number
# of GPUs. They are used for different parallel groups. sp_size is used for
# dit model and tp_size is used for encoder models.
fastvideo generate \
--model-path $MODEL_BASE \
--sp-size $num_gpus \
--tp-size $num_gpus \
--num-gpus $num_gpus \
--height 448 \
--width 832 \
--num-frames 77 \
--num-inference-steps 50 \
--fps 16 \
--guidance-scale 6.0 \
--flow-shift 8.0 \
--VSA-sparsity 0.9 \
--prompt "A beautiful woman in a red dress walking down a street" \
--negative-prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
--seed 1024 \
--output-path outputs_video_1.3B_VSA/sparsity_0.9/