Compare commits
48
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
48f9690c47 | ||
|
|
32df3bcca3 | ||
|
|
b8c81191b6 | ||
|
|
e506c4074f | ||
|
|
1b3914da84 | ||
|
|
f471fd3f02 | ||
|
|
2e6d5c5304 | ||
|
|
5db34184c7 | ||
|
|
37252bf62c | ||
|
|
e9263f7d2b | ||
|
|
93afb86c20 | ||
|
|
0d944ba9c1 | ||
|
|
3d78604281 | ||
|
|
0694b0c5eb | ||
|
|
b6c5644d40 | ||
|
|
a9089fa358 | ||
|
|
6079a98fd7 | ||
|
|
65c0fcb633 | ||
|
|
82e3641264 | ||
|
|
8741d204a5 | ||
|
|
cdc85f58a8 | ||
|
|
0262d2f089 | ||
|
|
62c0343465 | ||
|
|
1e1a023fb0 | ||
|
|
1d2517ad8e | ||
|
|
d41186cb4a | ||
|
|
78e0c7eec9 | ||
|
|
1c41a94b62 | ||
|
|
2e66aafe20 | ||
|
|
55074bda76 | ||
|
|
de65bec2b7 | ||
|
|
7664dd0de3 | ||
|
|
019a88ced4 | ||
|
|
72de11abcc | ||
|
|
d71a4ebffc | ||
|
|
1089ab43bf | ||
|
|
97d4b984c9 | ||
|
|
2a8953d74d | ||
|
|
8801b10da7 | ||
|
|
6b413f2ec4 | ||
|
|
28b72694aa | ||
|
|
4afb0cfe4f | ||
|
|
3eec1281cf | ||
|
|
0660489e38 | ||
|
|
dd871a17bf | ||
|
|
dc11529862 | ||
|
|
ffabf85e31 | ||
|
|
c0026ca5ba |
@@ -0,0 +1,66 @@
|
||||
env:
|
||||
IMAGE_VERSION: "py3.12-latest"
|
||||
|
||||
steps:
|
||||
- block: "Start Build"
|
||||
blocked_state: "running"
|
||||
prompt: "Approve build?"
|
||||
|
||||
- label: "Trigger Tests"
|
||||
command: |
|
||||
echo "Current working directory: $(pwd)"
|
||||
echo "Current branch:"
|
||||
git branch --show-current
|
||||
echo "Full diff:"
|
||||
git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD
|
||||
plugins:
|
||||
- monorepo-diff#v1.4.0:
|
||||
diff: "git diff --name-only $BUILDKITE_PULL_REQUEST_BASE_BRANCH...HEAD"
|
||||
watch:
|
||||
- path:
|
||||
- "fastvideo/v1/models/encoders/**"
|
||||
- "fastvideo/v1/models/loaders/**"
|
||||
- "fastvideo/v1/tests/encoders/**"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Encoder Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=encoder
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/models/vaes/**"
|
||||
- "fastvideo/v1/models/loaders/**"
|
||||
- "fastvideo/v1/tests/vaes/**"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "VAE Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=vae
|
||||
agents:
|
||||
queue: "default"
|
||||
- path:
|
||||
- "fastvideo/v1/models/dits/**"
|
||||
- "fastvideo/v1/models/loaders/**"
|
||||
- "fastvideo/v1/tests/transformers/**"
|
||||
- "fastvideo/v1/layers/**"
|
||||
- "fastvideo/v1/attention/**"
|
||||
config:
|
||||
command: "timeout 30m .buildkite/scripts/pr_test.sh"
|
||||
label: "Transformer Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=transformer
|
||||
agents:
|
||||
queue: "default"
|
||||
- path: "fastvideo/v1/**/*.py"
|
||||
config:
|
||||
command: "timeout 60m .buildkite/scripts/pr_test.sh"
|
||||
label: "SSIM Tests"
|
||||
env:
|
||||
- BUILDKITE_CLEAN_CHECKOUT=true
|
||||
- TEST_TYPE=ssim
|
||||
agents:
|
||||
queue: "default"
|
||||
Executable
+91
@@ -0,0 +1,91 @@
|
||||
#!/bin/bash
|
||||
set -uo pipefail
|
||||
|
||||
log() {
|
||||
echo "[$(date '+%Y-%m-%d %H:%M:%S')] $1"
|
||||
}
|
||||
|
||||
log "=== Starting Modal test execution ==="
|
||||
|
||||
# Change to the project directory
|
||||
cd "$(dirname "$0")/../.."
|
||||
PROJECT_ROOT=$(pwd)
|
||||
log "Project root: $PROJECT_ROOT"
|
||||
|
||||
# Install Modal if not available
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Modal not found, installing..."
|
||||
python3 -m pip install modal
|
||||
|
||||
# Verify installation
|
||||
if ! python3 -m modal --version &> /dev/null; then
|
||||
log "Error: Failed to install modal. Please install it manually."
|
||||
exit 1
|
||||
fi
|
||||
fi
|
||||
|
||||
log "modal version: $(python3 -m modal --version)"
|
||||
|
||||
# Set up Modal authentication using Buildkite secrets
|
||||
log "Setting up Modal authentication from Buildkite secrets..."
|
||||
MODAL_TOKEN_ID=$(buildkite-agent secret get modal_token_id)
|
||||
MODAL_TOKEN_SECRET=$(buildkite-agent secret get modal_token_secret)
|
||||
|
||||
if [ -n "$MODAL_TOKEN_ID" ] && [ -n "$MODAL_TOKEN_SECRET" ]; then
|
||||
log "Retrieved Modal credentials from Buildkite secrets"
|
||||
python3 -m modal token set --token-id "$MODAL_TOKEN_ID" --token-secret "$MODAL_TOKEN_SECRET" --profile buildkite-ci --activate --verify
|
||||
if [ $? -eq 0 ]; then
|
||||
log "Modal authentication successful"
|
||||
else
|
||||
log "Error: Failed to set Modal credentials"
|
||||
exit 1
|
||||
fi
|
||||
else
|
||||
log "Error: Could not retrieve Modal credentials from Buildkite secrets."
|
||||
log "Please ensure 'modal_token_id' and 'modal_token_secret' secrets are set in Buildkite."
|
||||
exit 1
|
||||
fi
|
||||
|
||||
MODAL_TEST_FILE="fastvideo/v1/tests/modal/pr_test.py"
|
||||
|
||||
if [ -z "${TEST_TYPE:-}" ]; then
|
||||
log "Error: TEST_TYPE environment variable is not set"
|
||||
exit 1
|
||||
fi
|
||||
log "Test type: $TEST_TYPE"
|
||||
|
||||
case "$TEST_TYPE" in
|
||||
"encoder")
|
||||
log "Running encoder tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_encoder_tests"
|
||||
;;
|
||||
"vae")
|
||||
log "Running VAE tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_vae_tests"
|
||||
;;
|
||||
"transformer")
|
||||
log "Running transformer tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_transformer_tests"
|
||||
;;
|
||||
"ssim")
|
||||
log "Running SSIM tests..."
|
||||
MODAL_COMMAND="python3 -m modal run $MODAL_TEST_FILE::run_ssim_tests"
|
||||
;;
|
||||
*)
|
||||
log "Error: Unknown test type: $TEST_TYPE"
|
||||
exit 1
|
||||
;;
|
||||
esac
|
||||
|
||||
log "Executing: $MODAL_COMMAND"
|
||||
eval "$MODAL_COMMAND"
|
||||
TEST_EXIT_CODE=$?
|
||||
|
||||
if [ $TEST_EXIT_CODE -eq 0 ]; then
|
||||
log "Modal test completed successfully"
|
||||
else
|
||||
log "Error: Modal test failed with exit code: $TEST_EXIT_CODE"
|
||||
fi
|
||||
|
||||
log "=== Test execution completed with exit code: $TEST_EXIT_CODE ==="
|
||||
exit $TEST_EXIT_CODE
|
||||
@@ -160,8 +160,7 @@ def execute_command(pod_id):
|
||||
setup_steps = [
|
||||
"tar -xzf /tmp/repo.tar.gz --no-same-owner -C /workspace/",
|
||||
f"cd /workspace/{repo_name}",
|
||||
"source /opt/conda/etc/profile.d/conda.sh",
|
||||
"conda activate fastvideo-dev",
|
||||
"source $HOME/.local/bin/env && source /opt/venv/bin/activate",
|
||||
args.test_command
|
||||
]
|
||||
|
||||
|
||||
@@ -12,12 +12,14 @@ on:
|
||||
paths:
|
||||
- "fastvideo/**/*.py"
|
||||
- ".github/workflows/pr-test.yml"
|
||||
- "pyproject.toml"
|
||||
- "docker/Dockerfile.python3.12"
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
custom_image:
|
||||
description: "Custom image from this repository (default: fastvideo-dev:latest)"
|
||||
description: "Custom image from this repository (default: fastvideo-dev:py3.12-latest)"
|
||||
required: false
|
||||
default: "fastvideo-dev:latest"
|
||||
default: "fastvideo-dev:py3.12-latest"
|
||||
type: string
|
||||
run_encoder_test:
|
||||
description: "Run encoder-test"
|
||||
@@ -39,6 +41,26 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_training_test:
|
||||
description: "Run training-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_training_test_VSA:
|
||||
description: "Run training-test-VSA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_inference_test_STA:
|
||||
description: "Run inference-test-STA"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_nightly_test:
|
||||
description: "Run nightly-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
|
||||
env:
|
||||
PYTHONUNBUFFERED: "1"
|
||||
@@ -59,6 +81,9 @@ jobs:
|
||||
encoder-test: ${{ steps.filter.outputs.encoder-test }}
|
||||
vae-test: ${{ steps.filter.outputs.vae-test }}
|
||||
transformer-test: ${{ steps.filter.outputs.transformer-test }}
|
||||
training-test: ${{ steps.filter.outputs.training-test }}
|
||||
training-test-VSA: ${{ steps.filter.outputs.training-test-VSA }}
|
||||
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
@@ -69,16 +94,34 @@ jobs:
|
||||
- 'fastvideo/v1/models/encoders/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/encoders/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
vae-test:
|
||||
- 'fastvideo/v1/models/vaes/**'
|
||||
- 'fastvideo/v1/models/loaders/**'
|
||||
- 'fastvideo/v1/tests/vaes/**'
|
||||
- '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/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
training-test:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
training-test-VSA:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
inference-test-STA:
|
||||
- 'fastvideo/v1/**'
|
||||
- 'pyproject.toml'
|
||||
- 'docker/Dockerfile.python3.12'
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
@@ -91,8 +134,8 @@ jobs:
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/encoders -s"
|
||||
image: "ghcr.io/${{ github.repository }}/${{ 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:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -109,8 +152,8 @@ jobs:
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/vaes -s"
|
||||
image: "ghcr.io/${{ github.repository }}/${{ 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:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -127,8 +170,8 @@ jobs:
|
||||
gpu_type: "NVIDIA L40S"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/transformers -s"
|
||||
image: "ghcr.io/${{ github.repository }}/${{ 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:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
@@ -155,11 +198,88 @@ jobs:
|
||||
volume_size: 200
|
||||
disk_size: 200
|
||||
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:${{ matrix.python-version.tag }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
|
||||
test_command: "uv pip install -e .[test] && pytest ./fastvideo/v1/tests/ssim -vs"
|
||||
timeout_minutes: 60
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
training-test:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && 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:
|
||||
job_id: "training-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
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:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
training-test-VSA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test_VSA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "training-test-VSA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ 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:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
inference-test-STA:
|
||||
needs: change-filter
|
||||
if: >-
|
||||
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_inference_test_STA == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "inference-test-STA"
|
||||
gpu_type: "NVIDIA H100 NVL"
|
||||
gpu_count: 1
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ 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 }}
|
||||
|
||||
nightly-test:
|
||||
if: >-
|
||||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
|
||||
uses: ./.github/workflows/runpod-test.yml
|
||||
with:
|
||||
job_id: "nightly-test"
|
||||
gpu_type: "NVIDIA A40"
|
||||
gpu_count: 4
|
||||
volume_size: 100
|
||||
disk_size: 100
|
||||
image: "ghcr.io/${{ github.repository }}/${{ 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:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
|
||||
runpod-cleanup:
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
|
||||
|
||||
@@ -43,6 +43,8 @@ on:
|
||||
required: true
|
||||
RUNPOD_PRIVATE_KEY:
|
||||
required: true
|
||||
WANDB_API_KEY:
|
||||
required: false
|
||||
|
||||
jobs:
|
||||
run-test:
|
||||
@@ -55,7 +57,7 @@ jobs:
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@v5
|
||||
with:
|
||||
python-version: "3.10"
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Set up SSH key
|
||||
run: |
|
||||
@@ -72,6 +74,7 @@ jobs:
|
||||
JOB_ID: ${{ inputs.job_id }}
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
GITHUB_RUN_ID: ${{ github.run_id }}
|
||||
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
|
||||
timeout-minutes: ${{ inputs.timeout_minutes }}
|
||||
run: >-
|
||||
python .github/scripts/runpod_api.py
|
||||
|
||||
@@ -91,7 +91,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
|
||||
|
||||
## Distillation and Finetuning
|
||||
- [Distillation Guide](https://hao-ai-lab.github.io/FastVideo/training/distillation.html)
|
||||
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetuning.html)
|
||||
- [Finetuning Guide](https://hao-ai-lab.github.io/FastVideo/training/finetune.html)
|
||||
|
||||
## 📑 Development Plan
|
||||
|
||||
@@ -111,7 +111,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
|
||||
|
||||
## 🤝 Contributing
|
||||
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/developer_guide/overview.html)
|
||||
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview.html)
|
||||
|
||||
## Acknowledgement
|
||||
We learned and reused code from the following projects:
|
||||
|
||||
@@ -6,6 +6,15 @@
|
||||
## Installation
|
||||
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
|
||||
sudo apt install gcc-11 g++-11
|
||||
|
||||
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
|
||||
## Environment Setup
|
||||
First, set up your CUDA environment:
|
||||
|
||||
@@ -154,6 +154,7 @@ def block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse
|
||||
return o, lse
|
||||
|
||||
def block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
grad_output = grad_output.contiguous()
|
||||
grad_q, grad_k, grad_v = block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
return grad_q, grad_k, grad_v
|
||||
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
FROM nvidia/cuda:12.4.1-devel-ubuntu20.04
|
||||
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
|
||||
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
|
||||
WORKDIR /FastVideo
|
||||
|
||||
RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
@@ -9,17 +11,25 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
git \
|
||||
ca-certificates \
|
||||
openssh-server \
|
||||
zsh \
|
||||
vim \
|
||||
curl \
|
||||
gcc-11 \
|
||||
g++-11 \
|
||||
clang-11 \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
RUN wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh && \
|
||||
bash Miniconda3-latest-Linux-x86_64.sh -b -p /opt/conda && \
|
||||
rm Miniconda3-latest-Linux-x86_64.sh
|
||||
# Set up C++20 compilers for ThunderKittens
|
||||
RUN update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
|
||||
|
||||
ENV PATH=/opt/conda/bin:$PATH
|
||||
# Set CUDA environment variables
|
||||
ENV CUDA_HOME=/usr/local/cuda-12.8
|
||||
ENV PATH=${CUDA_HOME}/bin:${PATH}
|
||||
ENV LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
|
||||
RUN conda create --name fastvideo-dev python=3.12.9 -y
|
||||
|
||||
SHELL ["/bin/bash", "-c"]
|
||||
# Install uv and source its environment
|
||||
RUN curl -LsSf https://astral.sh/uv/install.sh | sh && \
|
||||
echo 'source $HOME/.local/bin/env' >> /root/.bashrc
|
||||
|
||||
# Copy just the pyproject.toml first to leverage Docker cache
|
||||
COPY pyproject.toml ./
|
||||
@@ -27,22 +37,36 @@ COPY pyproject.toml ./
|
||||
# Create a dummy README to satisfy the installation
|
||||
RUN echo "# Placeholder" > README.md
|
||||
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir --upgrade pip && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir .[dev] && \
|
||||
conda run -n fastvideo-dev pip install --no-cache-dir flash-attn==2.7.4.post1 --no-build-isolation && \
|
||||
conda clean -afy
|
||||
# Create and activate virtual environment with specific Python version and seed
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
uv venv --python 3.12 --seed /opt/venv && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir --upgrade pip && \
|
||||
uv pip install --no-cache-dir .[dev] && \
|
||||
uv pip install --no-cache-dir flash-attn==2.8.0.post2 --no-build-isolation
|
||||
|
||||
COPY . .
|
||||
|
||||
RUN conda run -n fastvideo-dev pip install --no-cache-dir -e .[dev]
|
||||
# Install dependencies using uv and set up shell configuration
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
uv pip install --no-cache-dir -e .[dev] && \
|
||||
git config --unset-all http.https://github.com/.extraheader || true && \
|
||||
echo 'source /opt/venv/bin/activate' >> /root/.bashrc && \
|
||||
echo 'if [ -n "$ZSH_VERSION" ] && [ -f ~/.zshrc ]; then . ~/.zshrc; elif [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
|
||||
# Remove authentication headers
|
||||
RUN git config --unset-all http.https://github.com/.extraheader || true
|
||||
# Install STA (Sliding Tile Attention)
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_sta.py install
|
||||
|
||||
# Set up automatic conda environment activation for all shells
|
||||
RUN echo 'source /opt/conda/etc/profile.d/conda.sh' >> /root/.bashrc && \
|
||||
echo 'conda activate fastvideo-dev' >> /root/.bashrc && \
|
||||
# Ensure .bashrc is sourced for SSH login shells
|
||||
echo 'if [ -f ~/.bashrc ]; then . ~/.bashrc; fi' > /root/.profile
|
||||
# Install VSA
|
||||
RUN source $HOME/.local/bin/env && \
|
||||
source /opt/venv/bin/activate && \
|
||||
cd csrc/attn && \
|
||||
git submodule update --init --recursive && \
|
||||
python setup_vsa.py install
|
||||
|
||||
EXPOSE 22
|
||||
@@ -1,7 +1,7 @@
|
||||
(sta-demo)=
|
||||
|
||||
# 🔍 Demo
|
||||
There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
This is is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
|
||||
<div style="text-align: center;">
|
||||
<video controls width="800">
|
||||
@@ -9,3 +9,9 @@ There is a demo for 2D STA with window size (6,6) operating on a (10, 10) image.
|
||||
Your browser does not support the video tag.
|
||||
</video>
|
||||
</div>
|
||||
|
||||
You can run STA using the following command:
|
||||
|
||||
```bash
|
||||
bash scripts/inference/v1_inference_wan_STA.sh
|
||||
```
|
||||
|
||||
@@ -7,70 +7,40 @@ To save GPU memory, we precompute text embeddings and VAE latents to eliminate t
|
||||
We provide a sample dataset to help you get started. Download the source media using the following command:
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Image-Vid-Finetune-Src --local_dir=data/Image-Vid-Finetune-Src --repo_type=dataset
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/mini_i2v_dataset --local_dir=FastVideo/mini_i2v_dataset --repo_type=dataset
|
||||
```
|
||||
|
||||
The folder `crush-smol_raw/` contains raw videos and captions for testing preprocessing, while `crush-smol_preprocessed/` contains latents prepared for testing training.
|
||||
|
||||
To preprocess the dataset for fine-tuning or distillation, run:
|
||||
|
||||
```
|
||||
bash scripts/preprocess/preprocess_mochi_data.sh # for mochi
|
||||
bash scripts/preprocess/preprocess_hunyuan_data.sh # for hunyuan
|
||||
bash scripts/preprocess/v1_preprocess_wan_data_t2v # for wan
|
||||
```
|
||||
|
||||
The preprocessed dataset will be stored in `Image-Vid-Finetune-Mochi` or `Image-Vid-Finetune-HunYuan` correspondingly.
|
||||
|
||||
## Process your own dataset
|
||||
|
||||
If you wish to create your own dataset for finetuning or distillation, please structure you video dataset in the following format:
|
||||
If you wish to create your own dataset for finetuning or distillation, please refer `mini_i2v_dataset/crush-smol_raw/` to structure you video dataset in the following format:
|
||||
|
||||
```
|
||||
path_to_dataset_folder/
|
||||
├── media/
|
||||
│ ├── 0.jpg
|
||||
path_to_your_dataset_folder/
|
||||
├── videos/
|
||||
│ ├── 0.mp4
|
||||
│ ├── 1.mp4
|
||||
│ ├── 2.jpg
|
||||
├── video2caption.json
|
||||
└── merge.txt
|
||||
├── videos.txt
|
||||
└── prompt.txt
|
||||
```
|
||||
|
||||
Format the JSON file as a list, where each item represents a media source:
|
||||
To geranate the `videos2caption.json` and `merge.txt`, run
|
||||
|
||||
For image media,
|
||||
|
||||
```
|
||||
{
|
||||
"path": "0.jpg",
|
||||
"cap": ["captions"]
|
||||
}
|
||||
``` python
|
||||
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
|
||||
```
|
||||
|
||||
For video media,
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/v1_preprocess_****.sh` accordingly and run:
|
||||
|
||||
```
|
||||
{
|
||||
"path": "1.mp4",
|
||||
"resolution": {
|
||||
"width": 848,
|
||||
"height": 480
|
||||
},
|
||||
"fps": 30.0,
|
||||
"duration": 6.033333333333333,
|
||||
"cap": [
|
||||
"caption"
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
Use a txt file (merge.txt) to contain the source folder for media and the JSON file for meta information:
|
||||
|
||||
```
|
||||
path_to_media_source_foder,path_to_json_file
|
||||
```
|
||||
|
||||
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/preprocess_****_data.sh` accordingly and run:
|
||||
|
||||
```
|
||||
bash scripts/preprocess/preprocess_****_data.sh
|
||||
bash scripts/preprocess/v1_preprocess_****.sh
|
||||
```
|
||||
|
||||
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
|
||||
|
||||
@@ -16,6 +16,13 @@ bash scripts/finetune/finetune_mochi.sh # for mochi
|
||||
```
|
||||
|
||||
**Note that for finetuning, we did not tune the hyperparameters in the provided script.**
|
||||
## ⚡ Finetune with VSA
|
||||
Follow [data_preprocess.md](#v0-data-preprocess) to get parquet files for preproccessed latent, and then run:
|
||||
|
||||
```bash
|
||||
bash scripts/finetune/finetune_v1_VSA.sh
|
||||
```
|
||||
|
||||
## ⚡ Lora Finetune
|
||||
|
||||
Hunyuan supports Lora fine-tuning of videos up to 720p. Demos and prompts of Black-Myth-Wukong can be found in [here](https://huggingface.co/FastVideo/Hunyuan-Black-Myth-Wukong-lora-weight). You can download the Lora weight through:
|
||||
|
||||
@@ -5,7 +5,7 @@ export FASTVIDEO_ATTENTION_BACKEND=SLIDING_TILE_ATTN
|
||||
export MODEL_BASE=Wan-AI/Wan2.1-T2V-14B-Diffusers
|
||||
|
||||
base_port=29503
|
||||
num_gpu=$(nvidia-smi --query-gpu=gpu_name --format=csv,noheader | wc -l)
|
||||
num_gpu=1
|
||||
gpu_ids=$(seq 0 $((num_gpu-1)))
|
||||
skip_time_steps=12
|
||||
|
||||
@@ -14,7 +14,7 @@ STA_mode="STA_searching"
|
||||
for i in $gpu_ids; do
|
||||
port=$((base_port+i))
|
||||
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
|
||||
--prompt_path ./assets/prompt_extend_${i}.txt \
|
||||
--prompt_path ./assets/prompt_${i}.txt \
|
||||
--output_path $output_path \
|
||||
--STA_mode $STA_mode &
|
||||
sleep 1
|
||||
@@ -27,7 +27,7 @@ STA_mode="STA_tuning"
|
||||
for i in $gpu_ids; do
|
||||
port=$((base_port+i))
|
||||
CUDA_VISIBLE_DEVICES=$i MASTER_PORT=$port python examples/inference/sta_mask_search/wan_example.py \
|
||||
--prompt_path ./assets/prompt_extend_${i}.txt \
|
||||
--prompt_path ./assets/prompt_${i}.txt \
|
||||
--output_path $output_path \
|
||||
--STA_mode $STA_mode \
|
||||
--skip_time_steps $skip_time_steps &
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
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]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_i2v_finetune"
|
||||
--output_dir "$DATA_DIR/outputs/wan_i2v_finetune"
|
||||
--max_train_steps 5000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 8
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size 8
|
||||
--tp_size 8
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
--pretrained_model_name_or_path "Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_preprocessed_path "$VALIDATION_DIR"
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "40"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 6000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--cfg 0.0
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
)
|
||||
|
||||
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,97 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=FV_2N_14B
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --qos=hao
|
||||
#SBATCH --nodes=4
|
||||
#SBATCH --ntasks=4
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --nodelist=fs-mbz-gpu-[400-550]
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=4n_i2v/4n_i2v_%j.out
|
||||
#SBATCH --error=4n_i2v/4n_i2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_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
|
||||
|
||||
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
|
||||
srun torchrun \
|
||||
--nnodes $SLURM_JOB_NUM_NODES \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--node_rank $SLURM_PROCID \
|
||||
--rdzv_backend=c10d \
|
||||
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
|
||||
fastvideo/v1/training/wan_i2v_training_pipeline.py\
|
||||
--model_path Wan-AI/Wan2.1-I2V-14B-480P-Diffusers \
|
||||
--inference_mode False\
|
||||
--pretrained_model_name_or_path Wan-AI/Wan2.1-I2V-14B-480P-Diffusers \
|
||||
--cache_dir "/home/ray/.cache"\
|
||||
--data_path "$DATA_DIR"\
|
||||
--validation_preprocessed_path "$VALIDATION_DIR"\
|
||||
--train_batch_size=1\
|
||||
--num_latent_t 16 \
|
||||
--num_gpus $NUM_GPUS \
|
||||
--sp_size $NUM_GPUS \
|
||||
--tp_size $NUM_GPUS \
|
||||
--hsdp_replicate_dim $SLURM_JOB_NUM_NODES \
|
||||
--hsdp_shard_dim $NUM_GPUS \
|
||||
--train_sp_batch_size 1\
|
||||
--dataloader_num_workers 10\
|
||||
--gradient_accumulation_steps=2\
|
||||
--max_train_steps=10000 \
|
||||
--learning_rate=5e-5\
|
||||
--mixed_precision="bf16"\
|
||||
--checkpointing_steps=11000 \
|
||||
--validation_steps 100\
|
||||
--validation_sampling_steps "40" \
|
||||
--log_validation \
|
||||
--checkpoints_total_limit 3\
|
||||
--allow_tf32\
|
||||
--ema_start_step 0\
|
||||
--cfg 0.0\
|
||||
--output_dir="$DATA_DIR/outputs/wan_i2v_finetune_2n"\
|
||||
--tracker_project_name wan_i2v_finetune \
|
||||
--num_height 480 \
|
||||
--num_width 832 \
|
||||
--num_frames 77 \
|
||||
--validation_guidance_scale "1.0" \
|
||||
--num_euler_timesteps 50 \
|
||||
--multi_phased_distill_schedule "4000-1" \
|
||||
--weight_decay 1e-4 \
|
||||
--not_apply_cfg_solver \
|
||||
--dit_precision "fp32" \
|
||||
--max_grad_norm 1.0
|
||||
@@ -0,0 +1,24 @@
|
||||
# export WANDB_MODE="offline"
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_i2v/"
|
||||
VALIDATION_PATH="examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation.json"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 77 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--model_type $MODEL_TYPE \
|
||||
--train_fps 16 \
|
||||
--validation_dataset_file $VALIDATION_PATH \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--preprocess_task "i2v"
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
|
||||
"num_inference_steps": 40,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -0,0 +1,88 @@
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
|
||||
VALIDATION_DIR="data/crush-smol_processed_t2v/validation_parquet_dataset/"
|
||||
NUM_GPUS=4
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name "wan_t2v_finetune"
|
||||
--output_dir "outputs/wan_t2v_finetune"
|
||||
--max_train_steps 5000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 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 1 \
|
||||
--hsdp_shard_dim $NUM_GPUS \
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path $DATA_DIR
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_preprocessed_path $VALIDATION_DIR
|
||||
--validation_steps 100
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "1.0"
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--checkpointing_steps 6000
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--allow_tf32
|
||||
--checkpoints_total_limit 3
|
||||
--cfg 0.0
|
||||
--multi_phased_distill_schedule "4000-1"
|
||||
--not_apply_cfg_solver
|
||||
--dit_precision "fp32"
|
||||
--num_euler_timesteps 50
|
||||
--ema_start_step 0
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/v1/training/wan_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,98 @@
|
||||
#!/bin/bash
|
||||
#SBATCH --job-name=FV_2N_14B
|
||||
#SBATCH --partition=main
|
||||
#SBATCH --qos=hao
|
||||
#SBATCH --nodes=1
|
||||
#SBATCH --ntasks=1
|
||||
#SBATCH --ntasks-per-node=1
|
||||
#SBATCH --gres=gpu:8
|
||||
#SBATCH --cpus-per-task=128
|
||||
#SBATCH --nodelist=fs-mbz-gpu-[400-550]
|
||||
#SBATCH --mem=1440G
|
||||
#SBATCH --output=4n_i2v/4n_i2v_%j.out
|
||||
#SBATCH --error=4n_i2v/4n_i2v_%j.err
|
||||
#SBATCH --exclusive
|
||||
set -e -x
|
||||
|
||||
# Environment Setup
|
||||
source ~/conda/miniconda/bin/activate
|
||||
conda activate will-fv
|
||||
|
||||
# Basic Info
|
||||
export WANDB_MODE="online"
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export NODE_RANK=$SLURM_PROCID
|
||||
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
|
||||
export MASTER_ADDR=${nodes[0]}
|
||||
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_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]
|
||||
|
||||
# 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\
|
||||
--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
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,24 @@
|
||||
# export WANDB_MODE="offline"
|
||||
GPU_NUM=1 # 2,4,8
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_TYPE="wan"
|
||||
DATA_MERGE_PATH="data/crush-smol/merge.txt"
|
||||
OUTPUT_DIR="data/crush-smol_processed_t2v/"
|
||||
VALIDATION_PATH="examples/training/finetune/wan_t2v_1_3b/crush_smol/validation.json"
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
|
||||
--model_path $MODEL_PATH \
|
||||
--data_merge_path $DATA_MERGE_PATH \
|
||||
--preprocess_video_batch_size 8 \
|
||||
--max_height 480 \
|
||||
--max_width 832 \
|
||||
--num_frames 77 \
|
||||
--dataloader_num_workers 0 \
|
||||
--output_dir=$OUTPUT_DIR \
|
||||
--model_type $MODEL_TYPE \
|
||||
--train_fps 16 \
|
||||
--validation_dataset_file $VALIDATION_PATH \
|
||||
--samples_per_file 8 \
|
||||
--flush_frequency 8 \
|
||||
--preprocess_task "t2v"
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -6,6 +6,8 @@ from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
|
||||
from fastvideo.v1.utils import dict_to_3d_list
|
||||
|
||||
|
||||
def configure_sta(mode: str = 'STA_searching',
|
||||
layer_num: int = 40,
|
||||
@@ -349,21 +351,6 @@ def select_best_mask_strategy(
|
||||
return best_mask_strategy, overall_sparsity, strategy_counts
|
||||
|
||||
|
||||
def dict_to_3d_list(mask_strategy: Optional[Dict[str, List[int]]],
|
||||
t_max: int = 50,
|
||||
l_max: int = 60,
|
||||
h_max: int = 24) -> List[List[List[Optional[List[int]]]]]:
|
||||
result: List[List[List[Optional[List[int]]]]] = [[[
|
||||
None for _ in range(h_max)
|
||||
] for _ in range(l_max)] for _ in range(t_max)]
|
||||
if mask_strategy is None:
|
||||
return result
|
||||
for key, value in mask_strategy.items():
|
||||
t, layer_idx, h = map(int, key.split('_'))
|
||||
result[t][layer_idx][h] = value
|
||||
return result
|
||||
|
||||
|
||||
def save_mask_search_results(
|
||||
mask_search_final_result: List[Dict[str, List[float]]],
|
||||
prompt: str,
|
||||
|
||||
@@ -10,8 +10,8 @@ from fastvideo.v1.attention.selector import get_attn_backend
|
||||
|
||||
__all__ = [
|
||||
"DistributedAttention",
|
||||
"DistributedAttention_VSA",
|
||||
"LocalAttention",
|
||||
"DistributedAttention_VSA",
|
||||
"AttentionBackend",
|
||||
"AttentionMetadata",
|
||||
"AttentionMetadataBuilder",
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Type
|
||||
from typing import Any, List, Optional, Type
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
@@ -17,33 +17,11 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.utils import dict_to_3d_list
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
# TODO(will-refactor): move this to a utils file
|
||||
def dict_to_3d_list(
|
||||
mask_strategy: Dict[str,
|
||||
Any]) -> List[List[List[Optional[torch.Tensor]]]]:
|
||||
indices = [tuple(map(int, key.split('_'))) for key in mask_strategy]
|
||||
|
||||
max_timesteps_idx = max(
|
||||
timesteps_idx for timesteps_idx, layer_idx, head_idx in indices) + 1
|
||||
max_layer_idx = max(layer_idx
|
||||
for timesteps_idx, layer_idx, head_idx in indices) + 1
|
||||
max_head_idx = max(head_idx
|
||||
for timesteps_idx, layer_idx, head_idx in indices) + 1
|
||||
|
||||
result = [[[None for _ in range(max_head_idx)]
|
||||
for _ in range(max_layer_idx)] for _ in range(max_timesteps_idx)]
|
||||
|
||||
for key, value in mask_strategy.items():
|
||||
timesteps_idx, layer_idx, head_idx = map(int, key.split('_'))
|
||||
result[timesteps_idx][layer_idx][head_idx] = value
|
||||
|
||||
return result
|
||||
|
||||
|
||||
class RangeDict(dict):
|
||||
|
||||
def __getitem__(self, item: int) -> str:
|
||||
|
||||
@@ -5,7 +5,11 @@ from typing import List, Optional, Type
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from vsa import video_sparse_attn
|
||||
|
||||
try:
|
||||
from vsa import video_sparse_attn
|
||||
except ImportError:
|
||||
video_sparse_attn = None
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionImpl,
|
||||
@@ -68,14 +72,18 @@ class VideoSparseAttentionMetadataBuilder(AttentionMetadataBuilder):
|
||||
if forward_batch.latents is None:
|
||||
raise ValueError("latents cannot be None")
|
||||
|
||||
raw_latent_shape = forward_batch.latents.shape
|
||||
patch_size = fastvideo_args.dit_config.patch_size
|
||||
raw_latent_shape = forward_batch.raw_latent_shape
|
||||
if raw_latent_shape is None:
|
||||
raise ValueError("raw_latent_shape cannot be None")
|
||||
|
||||
patch_size = fastvideo_args.pipeline_config.dit_config.patch_size
|
||||
dit_seq_shape = [
|
||||
raw_latent_shape[2] // patch_size[0],
|
||||
raw_latent_shape[3] // patch_size[1],
|
||||
raw_latent_shape[4] // patch_size[2]
|
||||
]
|
||||
VSA_sparsity = forward_batch.VSA_sparsity
|
||||
|
||||
return VideoSparseAttentionMetadata(current_timestep=current_timestep,
|
||||
dit_seq_shape=dit_seq_shape,
|
||||
VSA_sparsity=VSA_sparsity)
|
||||
@@ -170,10 +178,15 @@ class VideoSparseAttentionImpl(AttentionImpl):
|
||||
value = value.transpose(1, 2).contiguous()
|
||||
gate_compress = gate_compress.transpose(1, 2).contiguous()
|
||||
|
||||
VSA_sparsity = attn_metadata.VSA_sparsity
|
||||
|
||||
cur_topk = math.ceil(
|
||||
(1 - attn_metadata.VSA_sparsity) *
|
||||
(1 - VSA_sparsity) *
|
||||
(self.img_seq_length / math.prod(self.VSA_base_tile_size)))
|
||||
|
||||
if video_sparse_attn is None:
|
||||
raise NotImplementedError("video_sparse_attn is not installed")
|
||||
|
||||
hidden_states = video_sparse_attn(
|
||||
query,
|
||||
key,
|
||||
|
||||
@@ -12,7 +12,7 @@ from fastvideo.v1.distributed.communication_op import (
|
||||
from fastvideo.v1.distributed.parallel_state import (get_sp_parallel_rank,
|
||||
get_sp_world_size)
|
||||
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
from fastvideo.v1.utils import get_compute_dtype
|
||||
|
||||
|
||||
@@ -26,8 +26,8 @@ class DistributedAttention(nn.Module):
|
||||
num_kv_heads: Optional[int] = None,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[
|
||||
AttentionBackendEnum, ...]] = None,
|
||||
prefix: str = "",
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
@@ -45,13 +45,13 @@ class DistributedAttention(nn.Module):
|
||||
dtype,
|
||||
supported_attention_backends=supported_attention_backends)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
causal=causal,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
prefix=f"{prefix}.impl",
|
||||
**extra_impl_args)
|
||||
self.attn_impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
causal=causal,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
prefix=f"{prefix}.impl",
|
||||
**extra_impl_args)
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
self.num_kv_heads = num_kv_heads
|
||||
@@ -100,7 +100,7 @@ class DistributedAttention(nn.Module):
|
||||
scatter_dim=2,
|
||||
gather_dim=1)
|
||||
# Apply backend-specific preprocess_qkv
|
||||
qkv = self.impl.preprocess_qkv(qkv, ctx_attn_metadata)
|
||||
qkv = self.attn_impl.preprocess_qkv(qkv, ctx_attn_metadata)
|
||||
|
||||
# Concatenate with replicated QKV if provided
|
||||
if replicated_q is not None:
|
||||
@@ -116,7 +116,7 @@ class DistributedAttention(nn.Module):
|
||||
|
||||
q, k, v = qkv.chunk(3, dim=0)
|
||||
|
||||
output = self.impl.forward(q, k, v, ctx_attn_metadata)
|
||||
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
||||
|
||||
# Redistribute back if using sequence parallelism
|
||||
replicated_output = None
|
||||
@@ -127,7 +127,7 @@ class DistributedAttention(nn.Module):
|
||||
replicated_output = sequence_model_parallel_all_gather(
|
||||
replicated_output.contiguous(), dim=2)
|
||||
# Apply backend-specific postprocess_output
|
||||
output = self.impl.postprocess_output(output, ctx_attn_metadata)
|
||||
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
|
||||
|
||||
output = sequence_model_parallel_all_to_all_4D(output,
|
||||
scatter_dim=1,
|
||||
@@ -183,18 +183,17 @@ class DistributedAttention_VSA(DistributedAttention):
|
||||
scatter_dim=2,
|
||||
gather_dim=1)
|
||||
|
||||
qkvg = self.impl.preprocess_qkv(
|
||||
qkvg, ctx_attn_metadata) # (yongqi) pass latent shape here?
|
||||
qkvg = self.attn_impl.preprocess_qkv(qkvg, ctx_attn_metadata)
|
||||
|
||||
q, k, v, gate_compress = qkvg.chunk(4, dim=0)
|
||||
output = self.impl.forward(q, k, v, gate_compress,
|
||||
ctx_attn_metadata) # type: ignore[call-arg]
|
||||
output = self.attn_impl.forward(
|
||||
q, k, v, gate_compress, ctx_attn_metadata) # type: ignore[call-arg]
|
||||
|
||||
# Redistribute back if using sequence parallelism
|
||||
replicated_output = None
|
||||
|
||||
# Apply backend-specific postprocess_output
|
||||
output = self.impl.postprocess_output(output, ctx_attn_metadata)
|
||||
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
|
||||
|
||||
output = sequence_model_parallel_all_to_all_4D(output,
|
||||
scatter_dim=1,
|
||||
@@ -212,8 +211,8 @@ class LocalAttention(nn.Module):
|
||||
num_kv_heads: Optional[int] = None,
|
||||
softmax_scale: Optional[float] = None,
|
||||
causal: bool = False,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[
|
||||
AttentionBackendEnum, ...]] = None,
|
||||
**extra_impl_args) -> None:
|
||||
super().__init__()
|
||||
if softmax_scale is None:
|
||||
@@ -229,12 +228,12 @@ class LocalAttention(nn.Module):
|
||||
dtype,
|
||||
supported_attention_backends=supported_attention_backends)
|
||||
impl_cls = attn_backend.get_impl_cls()
|
||||
self.impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
causal=causal,
|
||||
**extra_impl_args)
|
||||
self.attn_impl = impl_cls(num_heads=num_heads,
|
||||
head_size=head_size,
|
||||
softmax_scale=self.softmax_scale,
|
||||
num_kv_heads=num_kv_heads,
|
||||
causal=causal,
|
||||
**extra_impl_args)
|
||||
self.num_heads = num_heads
|
||||
self.head_size = head_size
|
||||
self.num_kv_heads = num_kv_heads
|
||||
@@ -265,5 +264,5 @@ class LocalAttention(nn.Module):
|
||||
forward_context: ForwardContext = get_forward_context()
|
||||
ctx_attn_metadata = forward_context.attn_metadata
|
||||
|
||||
output = self.impl.forward(q, k, v, ctx_attn_metadata)
|
||||
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
|
||||
return output
|
||||
|
||||
@@ -11,13 +11,13 @@ import torch
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.attention.backends.abstract import AttentionBackend
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.platforms import _Backend, current_platform
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum, current_platform
|
||||
from fastvideo.v1.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
|
||||
def backend_name_to_enum(backend_name: str) -> Optional[AttentionBackendEnum]:
|
||||
"""
|
||||
Convert a string backend name to a _Backend enum value.
|
||||
|
||||
@@ -27,11 +27,11 @@ def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
|
||||
loaded.
|
||||
"""
|
||||
assert backend_name is not None
|
||||
return _Backend[backend_name] if backend_name in _Backend.__members__ else \
|
||||
return AttentionBackendEnum[backend_name] if backend_name in AttentionBackendEnum.__members__ else \
|
||||
None
|
||||
|
||||
|
||||
def get_env_variable_attn_backend() -> Optional[_Backend]:
|
||||
def get_env_variable_attn_backend() -> Optional[AttentionBackendEnum]:
|
||||
'''
|
||||
Get the backend override specified by the FastVideo attention
|
||||
backend environment variable, if one is specified.
|
||||
@@ -53,10 +53,11 @@ def get_env_variable_attn_backend() -> Optional[_Backend]:
|
||||
#
|
||||
# THIS SELECTION TAKES PRECEDENCE OVER THE
|
||||
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
|
||||
forced_attn_backend: Optional[_Backend] = None
|
||||
forced_attn_backend: Optional[AttentionBackendEnum] = None
|
||||
|
||||
|
||||
def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
|
||||
def global_force_attn_backend(
|
||||
attn_backend: Optional[AttentionBackendEnum]) -> None:
|
||||
'''
|
||||
Force all attention operations to use a specified backend.
|
||||
|
||||
@@ -71,7 +72,7 @@ def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
|
||||
forced_attn_backend = attn_backend
|
||||
|
||||
|
||||
def get_global_forced_attn_backend() -> Optional[_Backend]:
|
||||
def get_global_forced_attn_backend() -> Optional[AttentionBackendEnum]:
|
||||
'''
|
||||
Get the currently-forced choice of attention backend,
|
||||
or None if auto-selection is currently enabled.
|
||||
@@ -82,7 +83,8 @@ def get_global_forced_attn_backend() -> Optional[_Backend]:
|
||||
def get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
) -> Type[AttentionBackend]:
|
||||
return _cached_get_attn_backend(head_size, dtype,
|
||||
supported_attention_backends)
|
||||
@@ -92,7 +94,8 @@ def get_attn_backend(
|
||||
def _cached_get_attn_backend(
|
||||
head_size: int,
|
||||
dtype: torch.dtype,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
) -> Type[AttentionBackend]:
|
||||
# Check whether a particular choice of backend was
|
||||
# previously forced.
|
||||
@@ -102,7 +105,7 @@ def _cached_get_attn_backend(
|
||||
if not supported_attention_backends:
|
||||
raise ValueError("supported_attention_backends is empty")
|
||||
selected_backend = None
|
||||
backend_by_global_setting: Optional[_Backend] = (
|
||||
backend_by_global_setting: Optional[AttentionBackendEnum] = (
|
||||
get_global_forced_attn_backend())
|
||||
if backend_by_global_setting is not None:
|
||||
selected_backend = backend_by_global_setting
|
||||
@@ -125,7 +128,7 @@ def _cached_get_attn_backend(
|
||||
|
||||
@contextmanager
|
||||
def global_force_attn_backend_context_manager(
|
||||
attn_backend: _Backend) -> Generator[None, None, None]:
|
||||
attn_backend: AttentionBackendEnum) -> Generator[None, None, None]:
|
||||
'''
|
||||
Globally force a FastVideo attention backend override within a
|
||||
context manager, reverting the global attention backend
|
||||
|
||||
@@ -4,7 +4,7 @@ from typing import Any, List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -13,12 +13,10 @@ class DiTArchConfig(ArchConfig):
|
||||
_compile_conditions: list = field(default_factory=list)
|
||||
_param_names_mapping: dict = field(default_factory=dict)
|
||||
_lora_param_names_mapping: dict = field(default_factory=dict)
|
||||
_supported_attention_backends: Tuple[_Backend,
|
||||
...] = (_Backend.SLIDING_TILE_ATTN,
|
||||
_Backend.SAGE_ATTN,
|
||||
_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA,
|
||||
_Backend.VIDEO_SPARSE_ATTN)
|
||||
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN)
|
||||
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
|
||||
@@ -6,14 +6,14 @@ import torch
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
@dataclass
|
||||
class EncoderArchConfig(ArchConfig):
|
||||
architectures: List[str] = field(default_factory=lambda: [])
|
||||
_supported_attention_backends: Tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA)
|
||||
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA)
|
||||
output_hidden_states: bool = False
|
||||
use_return_dict: bool = True
|
||||
|
||||
|
||||
@@ -1,4 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import argparse
|
||||
import dataclasses
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Union
|
||||
|
||||
@@ -129,3 +131,12 @@ class VAEConfig(ModelConfig):
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "VAEConfig":
|
||||
kwargs = {}
|
||||
for attr in dataclasses.fields(cls):
|
||||
value = getattr(args, attr.name, None)
|
||||
if value is not None:
|
||||
kwargs[attr.name] = value
|
||||
return cls(**kwargs)
|
||||
|
||||
@@ -3,7 +3,7 @@ from fastvideo.v1.configs.pipelines.base import (PipelineConfig,
|
||||
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
|
||||
HunyuanConfig)
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_for_name)
|
||||
get_pipeline_config_cls_from_name)
|
||||
from fastvideo.v1.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.v1.configs.pipelines.wan import (WanI2V480PConfig,
|
||||
WanI2V720PConfig,
|
||||
@@ -14,5 +14,5 @@ __all__ = [
|
||||
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
|
||||
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
|
||||
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
|
||||
"get_pipeline_config_cls_for_name"
|
||||
"get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -1,19 +1,31 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
from dataclasses import asdict, dataclass, field, fields
|
||||
from typing import Any, Callable, Dict, Optional, Tuple, cast
|
||||
from enum import Enum
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
|
||||
VAEConfig)
|
||||
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput
|
||||
from fastvideo.v1.configs.utils import update_config_from_args
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import shallow_asdict
|
||||
from fastvideo.v1.utils import (FlexibleArgumentParser, StoreBoolean,
|
||||
shallow_asdict)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class STA_Mode(str, Enum):
|
||||
"""STA (Sliding Tile Attention) modes."""
|
||||
STA_INFERENCE = "STA_inference"
|
||||
STA_SEARCHING = "STA_searching"
|
||||
STA_TUNING = "STA_tuning"
|
||||
STA_TUNING_CFG = "STA_tuning_cfg"
|
||||
NONE = None
|
||||
|
||||
|
||||
def preprocess_text(prompt: str) -> str:
|
||||
return prompt
|
||||
|
||||
@@ -22,59 +34,282 @@ def postprocess_text(output: BaseEncoderOutput) -> torch.tensor:
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
# config for a single pipeline
|
||||
@dataclass
|
||||
class PipelineConfig:
|
||||
"""Base configuration for all pipeline architectures."""
|
||||
model_path: str = ""
|
||||
pipeline_config_path: Optional[str] = None
|
||||
|
||||
# Video generation parameters
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: Optional[float] = None
|
||||
disable_autocast: bool = False
|
||||
|
||||
# Model configuration
|
||||
precision: str = "bf16"
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
dit_precision: str = "bf16"
|
||||
|
||||
# VAE configuration
|
||||
vae_config: VAEConfig = field(default_factory=VAEConfig)
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True
|
||||
vae_sp: bool = True
|
||||
vae_config: VAEConfig = field(default_factory=VAEConfig)
|
||||
|
||||
# DiT configuration
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
# Image encoder configuration
|
||||
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
|
||||
image_encoder_precision: str = "fp32"
|
||||
|
||||
# Text encoder configuration
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp16", ))
|
||||
DEFAULT_TEXT_ENCODER_PRECISIONS = ("fp16", )
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (EncoderConfig(), ))
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp16", ))
|
||||
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (preprocess_text, ))
|
||||
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(postprocess_text, ))
|
||||
|
||||
# STA (Spatial-Temporal Attention) parameters
|
||||
# LoRA parameters
|
||||
lora_path: Optional[str] = None
|
||||
lora_nickname: Optional[
|
||||
str] = "default" # for swapping adapters in the pipeline
|
||||
lora_target_names: Optional[List[
|
||||
str]] = None # can restrict list of layers to adapt, e.g. ["q_proj"]
|
||||
|
||||
# StepVideo specific parameters
|
||||
pos_magic: Optional[str] = None
|
||||
neg_magic: Optional[str] = None
|
||||
timesteps_scale: Optional[bool] = None
|
||||
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
STA_mode: str = "STA_inference"
|
||||
STA_mode: STA_Mode = STA_Mode.STA_INFERENCE
|
||||
skip_time_steps: int = 15
|
||||
|
||||
# Compilation
|
||||
enable_torch_compile: bool = False
|
||||
# enable_torch_compile: bool = False
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser,
|
||||
prefix: str = "") -> FlexibleArgumentParser:
|
||||
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
|
||||
|
||||
# model_path will be conflicting with the model_path in FastVideoArgs,
|
||||
# so we add it separately if prefix is not empty
|
||||
if prefix_with_dot != "":
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}model-path",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}model_path",
|
||||
default=PipelineConfig.model_path,
|
||||
help="Path to the pretrained model",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}pipeline-config-path",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}pipeline_config_path",
|
||||
default=PipelineConfig.pipeline_config_path,
|
||||
help="Path to the pipeline config",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}embedded-cfg-scale",
|
||||
type=float,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}embedded_cfg_scale",
|
||||
default=PipelineConfig.embedded_cfg_scale,
|
||||
help="Embedded CFG scale",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}flow-shift",
|
||||
type=float,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}flow_shift",
|
||||
default=PipelineConfig.flow_shift,
|
||||
help="Flow shift parameter",
|
||||
)
|
||||
|
||||
# DiT configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}dit-precision",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}dit_precision",
|
||||
default=PipelineConfig.dit_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for the DiT model",
|
||||
)
|
||||
|
||||
# VAE configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}vae-precision",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}vae_precision",
|
||||
default=PipelineConfig.vae_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for VAE",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}vae-tiling",
|
||||
action=StoreBoolean,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}vae_tiling",
|
||||
default=PipelineConfig.vae_tiling,
|
||||
help="Enable VAE tiling",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}vae-sp",
|
||||
action=StoreBoolean,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}vae_sp",
|
||||
help="Enable VAE spatial parallelism",
|
||||
)
|
||||
|
||||
# Text encoder configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}text-encoder-precisions",
|
||||
nargs="+",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}text_encoder_precisions",
|
||||
default=PipelineConfig.DEFAULT_TEXT_ENCODER_PRECISIONS,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for each text encoder",
|
||||
)
|
||||
|
||||
# Image encoder configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}image-encoder-precision",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}image_encoder_precision",
|
||||
default=PipelineConfig.image_encoder_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for image encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}pos_magic",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}pos_magic",
|
||||
default=PipelineConfig.pos_magic,
|
||||
help="Positive magic prompt for sampling, used in stepvideo",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}neg_magic",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}neg_magic",
|
||||
default=PipelineConfig.neg_magic,
|
||||
help="Negative magic prompt for sampling, used in stepvideo",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}timesteps_scale",
|
||||
type=bool,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}timesteps_scale",
|
||||
default=PipelineConfig.timesteps_scale,
|
||||
help=
|
||||
"Bool for applying scheduler scale in set_timesteps, used in stepvideo",
|
||||
)
|
||||
|
||||
# Add VAE configuration arguments
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEConfig
|
||||
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
|
||||
|
||||
# Add DiT configuration arguments
|
||||
from fastvideo.v1.configs.models.dits.base import DiTConfig
|
||||
DiTConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}dit-config")
|
||||
|
||||
return parser
|
||||
|
||||
def update_config_from_dict(self,
|
||||
args: Dict[str, Any],
|
||||
prefix: str = "") -> None:
|
||||
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
|
||||
update_config_from_args(self, args, prefix, pop_args=True)
|
||||
update_config_from_args(self.vae_config,
|
||||
args,
|
||||
f"{prefix_with_dot}vae_config",
|
||||
pop_args=True)
|
||||
update_config_from_args(self.dit_config,
|
||||
args,
|
||||
f"{prefix_with_dot}dit_config",
|
||||
pop_args=True)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str) -> "PipelineConfig":
|
||||
"""
|
||||
use the pipeline class setting from model_path to match the pipeline config
|
||||
"""
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_for_name)
|
||||
pipeline_config_cls = get_pipeline_config_cls_for_name(model_path)
|
||||
if pipeline_config_cls is not None:
|
||||
pipeline_config = pipeline_config_cls()
|
||||
else:
|
||||
get_pipeline_config_cls_from_name)
|
||||
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
|
||||
|
||||
return cast(PipelineConfig, pipeline_config_cls(model_path=model_path))
|
||||
|
||||
@classmethod
|
||||
def from_kwargs(cls,
|
||||
kwargs: Dict[str, Any],
|
||||
config_cli_prefix: str = "") -> "PipelineConfig":
|
||||
"""
|
||||
Load PipelineConfig from kwargs Dictionary.
|
||||
kwargs: dictionary of kwargs
|
||||
config_cli_prefix: prefix of CLI arguments for this PipelineConfig instance
|
||||
"""
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
|
||||
prefix_with_dot = f"{config_cli_prefix}." if (config_cli_prefix.strip()
|
||||
!= "") else ""
|
||||
model_path: Optional[str] = kwargs.get(prefix_with_dot + 'model_path',
|
||||
None) or kwargs.get('model_path')
|
||||
pipeline_config_or_path: Optional[Union[str, PipelineConfig, Dict[
|
||||
str, Any]]] = kwargs.get(prefix_with_dot + 'pipeline_config',
|
||||
None) or kwargs.get('pipeline_config')
|
||||
if model_path is None:
|
||||
raise ValueError("model_path is required in kwargs")
|
||||
|
||||
# 1. Get the pipeline config class from the registry
|
||||
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
|
||||
|
||||
# 2. Instantiate PipelineConfig
|
||||
if pipeline_config_cls is None:
|
||||
logger.warning(
|
||||
"Couldn't find an optimal sampling param for %s. Using the default sampling param.",
|
||||
"Couldn't find pipeline config for %s. Using the default pipeline config.",
|
||||
model_path)
|
||||
pipeline_config = cls()
|
||||
else:
|
||||
pipeline_config = pipeline_config_cls()
|
||||
|
||||
return cast(PipelineConfig, pipeline_config)
|
||||
# 3. Load PipelineConfig from a json file or a PipelineConfig object if provided
|
||||
if isinstance(pipeline_config_or_path, str):
|
||||
pipeline_config.load_from_json(pipeline_config_or_path)
|
||||
kwargs[prefix_with_dot +
|
||||
'pipeline_config_path'] = pipeline_config_or_path
|
||||
elif isinstance(pipeline_config_or_path, PipelineConfig):
|
||||
pipeline_config = pipeline_config_or_path
|
||||
elif isinstance(pipeline_config_or_path, dict):
|
||||
pipeline_config.update_pipeline_config(pipeline_config_or_path)
|
||||
|
||||
# 4. Update PipelineConfig from CLI arguments if provided
|
||||
kwargs[prefix_with_dot + 'model_path'] = model_path
|
||||
pipeline_config.update_config_from_dict(kwargs, config_cli_prefix)
|
||||
return pipeline_config
|
||||
|
||||
def check_pipeline_config(self) -> None:
|
||||
if self.vae_sp and not self.vae_tiling:
|
||||
raise ValueError(
|
||||
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
|
||||
)
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.text_encoder_precisions):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
|
||||
)
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
)
|
||||
|
||||
if len(self.preprocess_text_funcs) != len(self.postprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
)
|
||||
|
||||
def dump_to_json(self, file_path: str):
|
||||
output_dict = shallow_asdict(self)
|
||||
|
||||
@@ -80,7 +80,7 @@ class HunyuanConfig(PipelineConfig):
|
||||
(llama_postprocess_text, clip_postprocess_text))
|
||||
|
||||
# Precision for each component
|
||||
precision: str = "bf16"
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "fp16"
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp16", "fp16"))
|
||||
|
||||
@@ -19,7 +19,7 @@ from fastvideo.v1.utils import (maybe_download_model_index,
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Registry maps specific model weights to their config classes
|
||||
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[PipelineConfig]] = {
|
||||
PIPE_NAME_TO_CONFIG: Dict[str, Type[PipelineConfig]] = {
|
||||
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
|
||||
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
|
||||
@@ -51,37 +51,74 @@ PIPELINE_FALLBACK_CONFIG: Dict[str, Type[PipelineConfig]] = {
|
||||
}
|
||||
|
||||
|
||||
def get_pipeline_config_cls_for_name(
|
||||
pipeline_name_or_path: str) -> Optional[type[PipelineConfig]]:
|
||||
"""Get the appropriate config class for specific pretrained weights."""
|
||||
def get_pipeline_config_cls_from_name(
|
||||
pipeline_name_or_path: str) -> Type[PipelineConfig]:
|
||||
"""Get the appropriate configuration class for a given pipeline name or path.
|
||||
|
||||
if os.path.exists(pipeline_name_or_path):
|
||||
config = verify_model_config_and_directory(pipeline_name_or_path)
|
||||
logger.warning(
|
||||
"FastVideo may not correctly identify the optimal config for this model, as the local directory may have been renamed."
|
||||
)
|
||||
else:
|
||||
config = maybe_download_model_index(pipeline_name_or_path)
|
||||
This function implements a multi-step lookup process to find the most suitable
|
||||
configuration class for a given pipeline. It follows this order:
|
||||
1. Exact match in the PIPE_NAME_TO_CONFIG
|
||||
2. Partial match in the PIPE_NAME_TO_CONFIG
|
||||
3. Fallback to class name in the model_index.json
|
||||
4. else raise an error
|
||||
|
||||
pipeline_name = config["_class_name"]
|
||||
Args:
|
||||
pipeline_name_or_path (str): The name or path of the pipeline. This can be:
|
||||
- A registered model ID (e.g., "FastVideo/FastHunyuan-diffusers")
|
||||
- A local path to a model directory
|
||||
- A model ID that will be downloaded
|
||||
|
||||
Returns:
|
||||
Type[PipelineConfig]: The configuration class that best matches the pipeline.
|
||||
This will be one of:
|
||||
- A specific weight configuration class if an exact match is found
|
||||
- A fallback configuration class based on the pipeline architecture
|
||||
- The base PipelineConfig class if no matches are found
|
||||
|
||||
Note:
|
||||
- For local paths, the function will verify the model configuration
|
||||
- For remote models, it will attempt to download the model index
|
||||
- Warning messages are logged when falling back to less specific configurations
|
||||
"""
|
||||
|
||||
pipeline_config_cls: Optional[Type[PipelineConfig]] = None
|
||||
|
||||
# First try exact match for specific weights
|
||||
if pipeline_name_or_path in WEIGHT_CONFIG_REGISTRY:
|
||||
return WEIGHT_CONFIG_REGISTRY[pipeline_name_or_path]
|
||||
if pipeline_name_or_path in PIPE_NAME_TO_CONFIG:
|
||||
pipeline_config_cls = PIPE_NAME_TO_CONFIG[pipeline_name_or_path]
|
||||
|
||||
# Try partial matches (for local paths that might include the weight ID)
|
||||
for registered_id, config_class in WEIGHT_CONFIG_REGISTRY.items():
|
||||
for registered_id, config_class in PIPE_NAME_TO_CONFIG.items():
|
||||
if registered_id in pipeline_name_or_path:
|
||||
return config_class
|
||||
|
||||
# If no match, try to use the fallback config
|
||||
fallback_config = None
|
||||
# Try to determine pipeline architecture for fallback
|
||||
for pipeline_type, detector in PIPELINE_DETECTOR.items():
|
||||
if detector(pipeline_name.lower()):
|
||||
fallback_config = PIPELINE_FALLBACK_CONFIG.get(pipeline_type)
|
||||
pipeline_config_cls = config_class
|
||||
break
|
||||
|
||||
logger.warning("No match found for pipeline %s, using fallback config %s.",
|
||||
pipeline_name_or_path, fallback_config)
|
||||
return fallback_config
|
||||
# If no match, try to use the fallback config
|
||||
if pipeline_config_cls is None:
|
||||
if os.path.exists(pipeline_name_or_path):
|
||||
config = verify_model_config_and_directory(pipeline_name_or_path)
|
||||
else:
|
||||
config = maybe_download_model_index(pipeline_name_or_path)
|
||||
logger.warning(
|
||||
"Trying to use the config from the model_index.json. FastVideo may not correctly identify the optimal config for this model in this situation."
|
||||
)
|
||||
|
||||
pipeline_name = config["_class_name"]
|
||||
# Try to determine pipeline architecture for fallback
|
||||
for pipeline_type, detector in PIPELINE_DETECTOR.items():
|
||||
if detector(pipeline_name.lower()):
|
||||
pipeline_config_cls = PIPELINE_FALLBACK_CONFIG.get(
|
||||
pipeline_type)
|
||||
break
|
||||
|
||||
if pipeline_config_cls is not None:
|
||||
logger.warning(
|
||||
"No match found for pipeline %s, using fallback config %s.",
|
||||
pipeline_name_or_path, pipeline_config_cls)
|
||||
|
||||
if pipeline_config_cls is None:
|
||||
raise ValueError(
|
||||
f"No match found for pipeline {pipeline_name_or_path}, please check the pipeline name or path."
|
||||
)
|
||||
|
||||
return pipeline_config_cls
|
||||
|
||||
@@ -39,7 +39,6 @@ class SamplingParam:
|
||||
num_inference_steps: int = 50
|
||||
guidance_scale: float = 1.0
|
||||
guidance_rescale: float = 0.0
|
||||
VSA_sparsity: float = 0.0
|
||||
|
||||
# TeaCache parameters
|
||||
enable_teacache: bool = False
|
||||
@@ -185,12 +184,6 @@ class SamplingParam:
|
||||
default=SamplingParam.image_path,
|
||||
help="Path to input image for image-to-video generation",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--VSA-sparsity",
|
||||
type=float,
|
||||
default=SamplingParam.VSA_sparsity,
|
||||
help="VSA attention sparsity",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
from typing import Any, Dict
|
||||
|
||||
|
||||
def update_config_from_args(config: Any,
|
||||
args_dict: Dict[str, Any],
|
||||
prefix: str = "",
|
||||
pop_args: bool = False) -> None:
|
||||
"""
|
||||
Update configuration object from arguments dictionary.
|
||||
|
||||
Args:
|
||||
config: The configuration object to update
|
||||
args_dict: Dictionary containing arguments
|
||||
prefix: Prefix for the configuration parameters in the args_dict.
|
||||
If None, assumes direct attribute mapping without prefix.
|
||||
"""
|
||||
# Handle top-level attributes (no prefix)
|
||||
args_not_to_remove = [
|
||||
'model_path',
|
||||
]
|
||||
args_to_remove = []
|
||||
if prefix.strip() == "":
|
||||
for key, value in args_dict.items():
|
||||
if hasattr(config, key) and value is not None:
|
||||
if key == "text_encoder_precisions" and isinstance(value, list):
|
||||
setattr(config, key, tuple(value))
|
||||
else:
|
||||
setattr(config, key, value)
|
||||
if pop_args:
|
||||
args_to_remove.append(key)
|
||||
else:
|
||||
# Handle nested attributes with prefix
|
||||
prefix_with_dot = f"{prefix}."
|
||||
for key, value in args_dict.items():
|
||||
if key.startswith(prefix_with_dot) and value is not None:
|
||||
attr_name = key[len(prefix_with_dot):]
|
||||
if hasattr(config, attr_name):
|
||||
setattr(config, attr_name, value)
|
||||
if pop_args:
|
||||
args_to_remove.append(key)
|
||||
|
||||
if pop_args:
|
||||
for key in args_to_remove:
|
||||
if key not in args_not_to_remove:
|
||||
args_dict.pop(key)
|
||||
@@ -1,19 +1,17 @@
|
||||
import os
|
||||
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from torchvision import transforms
|
||||
from torchvision.transforms import Lambda
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.v1.dataset.t2v_datasets import T2V_dataset
|
||||
from fastvideo.v1.dataset.parquet_dataset_map_style import (
|
||||
build_parquet_map_style_dataloader)
|
||||
from fastvideo.v1.dataset.preprocessing_datasets import (
|
||||
VideoCaptionMergedDataset)
|
||||
from fastvideo.v1.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
|
||||
from .parquet_dataset_map_style import build_parquet_map_style_dataloader
|
||||
|
||||
__all__ = ["build_parquet_map_style_dataloader"]
|
||||
from fastvideo.v1.dataset.validation_dataset import ValidationDataset
|
||||
|
||||
|
||||
def getdataset(args, start_idx=0) -> T2V_dataset:
|
||||
def getdataset(args) -> VideoCaptionMergedDataset:
|
||||
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
|
||||
norm_fun = Lambda(lambda x: 2.0 * x - 1.0)
|
||||
resize_topcrop = [
|
||||
@@ -31,15 +29,14 @@ def getdataset(args, start_idx=0) -> T2V_dataset:
|
||||
*resize_topcrop,
|
||||
norm_fun,
|
||||
])
|
||||
tokenizer_path = os.path.join(args.model_path, "tokenizer")
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
cache_dir=args.cache_dir)
|
||||
if args.dataset == "t2v":
|
||||
return T2V_dataset(args,
|
||||
transform=transform,
|
||||
temporal_sample=temporal_sample,
|
||||
tokenizer=tokenizer,
|
||||
transform_topcrop=transform_topcrop,
|
||||
start_idx=start_idx)
|
||||
return VideoCaptionMergedDataset(data_merge_path=args.data_merge_path,
|
||||
args=args,
|
||||
transform=transform,
|
||||
temporal_sample=temporal_sample,
|
||||
transform_topcrop=transform_topcrop)
|
||||
|
||||
raise NotImplementedError(args.dataset)
|
||||
|
||||
__all__ = [
|
||||
"build_parquet_map_style_dataloader", "ValidationDataset",
|
||||
"VideoCaptionMergedDataset"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,185 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import argparse
|
||||
import os
|
||||
import pathlib
|
||||
import time
|
||||
|
||||
import torch.distributed as dist
|
||||
import torch.distributed.checkpoint as dist_cp
|
||||
|
||||
from fastvideo.v1.dataset.parquet_dataset_iterable_style import (
|
||||
build_parquet_iterable_style_dataloader)
|
||||
from fastvideo.v1.distributed import get_world_rank
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_torch_device,
|
||||
maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Benchmark parquet iterable style dataset loading speed")
|
||||
parser.add_argument(
|
||||
"--path",
|
||||
type=str,
|
||||
help="Path to parquet dataset",
|
||||
)
|
||||
parser.add_argument("--batch_size",
|
||||
type=int,
|
||||
default=4,
|
||||
help="Batch size for DataLoader")
|
||||
parser.add_argument("--num_data_workers",
|
||||
type=int,
|
||||
help="Number of DataLoader workers")
|
||||
parser.add_argument("--num_epoch",
|
||||
type=int,
|
||||
default=2,
|
||||
help="Number of epoches to benchmark")
|
||||
parser.add_argument("--verify_resume",
|
||||
action="store_true",
|
||||
help="Verify resume")
|
||||
parser.add_argument(
|
||||
"--num_batches_per_epoch",
|
||||
type=int,
|
||||
default=1000,
|
||||
help="Number of batches to benchmark",
|
||||
)
|
||||
parser.add_argument('--checkpoint_path',
|
||||
type=str,
|
||||
default='dataloader_checkpoint',
|
||||
help='Path to save/load checkpoint')
|
||||
'''
|
||||
example launch command:
|
||||
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 2 --num_epoch 2 --num_batches_per_epoch 2 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path /mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents/ --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 100
|
||||
'''
|
||||
args = parser.parse_args()
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
maybe_init_distributed_environment_and_model_parallel(
|
||||
tp_size=(world_size + 1) // 2, sp_size=(world_size + 1) // 2)
|
||||
logger.info("Initialized distributed environment with world_size=%d",
|
||||
world_size)
|
||||
|
||||
# Create DataLoader with proper settings
|
||||
dataset, dataloader = build_parquet_iterable_style_dataloader(
|
||||
args.path, args.batch_size, args.num_data_workers)
|
||||
logger.info("Initialized dataloader")
|
||||
|
||||
if args.verify_resume:
|
||||
# First pass - record latent sums
|
||||
first_pass_sums = []
|
||||
for i, (latents, embeddings, masks,
|
||||
caption_text) in enumerate(dataloader):
|
||||
latent_sum = latents.sum().item()
|
||||
first_pass_sums.append(latent_sum)
|
||||
logger.info("Batch %d latent sum: %f", i, latent_sum)
|
||||
if i >= args.num_batches_per_epoch - 1:
|
||||
break
|
||||
|
||||
# Save dataloader state using distributed checkpoint
|
||||
checkpoint_dir = pathlib.Path(args.checkpoint_path)
|
||||
logger.info("Rank %d: Saving dataloader state to %s", get_world_rank(),
|
||||
checkpoint_dir)
|
||||
states = {"dataloader": dataloader}
|
||||
|
||||
begin_time = time.monotonic()
|
||||
dist_cp.save(states, checkpoint_id=checkpoint_dir.as_posix())
|
||||
end_time = time.monotonic()
|
||||
|
||||
logger.info("Rank %d: Saved checkpoint in %.2f seconds",
|
||||
get_world_rank(), end_time - begin_time)
|
||||
|
||||
# Make sure all processes wait for checkpoint to be saved
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
# Recreate dataloader and load state
|
||||
dataset, dataloader = build_parquet_iterable_style_dataloader(
|
||||
args.path, args.batch_size, args.num_data_workers)
|
||||
load_states = {"dataloader": dataloader}
|
||||
dist_cp.load(load_states, checkpoint_id=checkpoint_dir.as_posix())
|
||||
logger.info("Rank %d: Loaded dataloader state from %s",
|
||||
get_world_rank(), checkpoint_dir)
|
||||
|
||||
# Second pass - verify latent sums match
|
||||
for i, (latents, embeddings, masks) in enumerate(dataloader):
|
||||
latent_sum = latents.sum().item()
|
||||
first_pass_sums.append(latent_sum)
|
||||
logger.info("Batch %d latent sum: %f",
|
||||
i + args.num_batches_per_epoch, latent_sum)
|
||||
if i >= args.num_batches_per_epoch - 1:
|
||||
break
|
||||
|
||||
dataset, dataloader = build_parquet_iterable_style_dataloader(
|
||||
args.path, args.batch_size, args.num_data_workers)
|
||||
# Second pass - verify latent sums match
|
||||
second_pass_sums = []
|
||||
for i, (latents, embeddings, masks,
|
||||
caption_text) in enumerate(dataloader):
|
||||
latent_sum = latents.sum().item()
|
||||
second_pass_sums.append(latent_sum)
|
||||
logger.info("Batch %d latent sum: %f (should match first pass: %f)",
|
||||
i, latent_sum, first_pass_sums[i])
|
||||
if i >= args.num_batches_per_epoch * 2 - 1:
|
||||
break
|
||||
|
||||
# Verify all sums match
|
||||
if all(
|
||||
abs(a - b) < 1e-6
|
||||
for a, b in zip(first_pass_sums, second_pass_sums)):
|
||||
logger.info(
|
||||
"All latent sums match between passes - resume verification successful!"
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"Latent sums do not match between passes - resume verification failed!"
|
||||
)
|
||||
|
||||
start_time = time.time()
|
||||
total_samples = 0
|
||||
total_batches = 0
|
||||
for _ in range(args.num_epoch):
|
||||
for i, (latents, embeddings, masks,
|
||||
caption_text) in enumerate(dataloader):
|
||||
if i >= args.num_batches_per_epoch:
|
||||
break
|
||||
|
||||
# Move data to device
|
||||
latents = latents.to(get_torch_device())
|
||||
embeddings = embeddings.to(get_torch_device())
|
||||
|
||||
# Calculate actual batch size
|
||||
batch_size = latents.size(0)
|
||||
total_samples += batch_size
|
||||
total_batches += 1
|
||||
|
||||
# Print progress only from rank 0
|
||||
if get_world_rank() == 0 and (i + 1) % 10 == 0:
|
||||
elapsed = time.time() - start_time
|
||||
samples_per_sec = total_samples / elapsed
|
||||
logger.info("Batch %d/%d, Speed: %.2f samples/sec", i + 1,
|
||||
args.num_batches_per_epoch, samples_per_sec)
|
||||
|
||||
# Final statistics
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
if get_world_rank() == 0:
|
||||
elapsed = time.time() - start_time
|
||||
samples_per_sec = total_samples / elapsed
|
||||
|
||||
logger.info("\nBenchmark Results:")
|
||||
logger.info("Total time: %.2f seconds", elapsed)
|
||||
logger.info("Total samples: %d", total_samples)
|
||||
logger.info("Average speed: %.2f samples/sec", samples_per_sec)
|
||||
logger.info("Time per batch: %.2f ms", elapsed / total_batches * 1000)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
try:
|
||||
main()
|
||||
finally:
|
||||
cleanup_dist_env_and_memory()
|
||||
@@ -54,9 +54,9 @@ def main() -> None:
|
||||
help='Path to save/load checkpoint')
|
||||
'''
|
||||
example launch command:
|
||||
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 2 --verify_resume
|
||||
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 3 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5 --verify_resume
|
||||
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path /mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents/ --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 100
|
||||
'''
|
||||
args = parser.parse_args()
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
@@ -66,14 +66,18 @@ def main() -> None:
|
||||
world_size)
|
||||
|
||||
# Create DataLoader with proper settings
|
||||
dataloader = build_parquet_map_style_dataloader(args.path, args.batch_size,
|
||||
args.num_data_workers)
|
||||
dataset, dataloader = build_parquet_map_style_dataloader(
|
||||
args.path, args.batch_size, args.num_data_workers)
|
||||
logger.info("Initialized dataloader with %d batches", len(dataloader))
|
||||
|
||||
if args.verify_resume:
|
||||
# First pass - record latent sums
|
||||
first_pass_sums = []
|
||||
for i, (latents, embeddings, masks,
|
||||
data_indices) in enumerate(dataloader):
|
||||
logger.info("Batch %d data_indices: %s", i, data_indices)
|
||||
caption_text) in enumerate(dataloader):
|
||||
latent_sum = latents.sum().item()
|
||||
first_pass_sums.append(latent_sum)
|
||||
logger.info("Batch %d latent sum: %f", i, latent_sum)
|
||||
if i >= args.num_batches_per_epoch - 1:
|
||||
break
|
||||
|
||||
@@ -94,41 +98,54 @@ def main() -> None:
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
dataloader = build_parquet_map_style_dataloader(args.path,
|
||||
args.batch_size,
|
||||
args.num_data_workers)
|
||||
# Load dataloader state using distributed checkpoint
|
||||
logger.info("Rank %d: Loading dataloader state from %s",
|
||||
get_world_rank(), checkpoint_dir)
|
||||
# Recreate dataloader and load state
|
||||
dataset, dataloader = build_parquet_map_style_dataloader(
|
||||
args.path, args.batch_size, args.num_data_workers)
|
||||
load_states = {"dataloader": dataloader}
|
||||
dist_cp.load(load_states, checkpoint_id=checkpoint_dir.as_posix())
|
||||
logger.info("Rank %d: Loaded dataloader state from %s",
|
||||
get_world_rank(), checkpoint_dir)
|
||||
|
||||
for i, (latents, embeddings, masks,
|
||||
data_indices) in enumerate(dataloader):
|
||||
logger.info("Batch %d data_indices: %s", i, data_indices)
|
||||
caption_text) in enumerate(dataloader):
|
||||
latent_sum = latents.sum().item()
|
||||
first_pass_sums.append(latent_sum)
|
||||
logger.info("Batch %d latent sum: %f",
|
||||
i + args.num_batches_per_epoch, latent_sum)
|
||||
if i >= args.num_batches_per_epoch - 1:
|
||||
break
|
||||
|
||||
logger.info("Restart from the beginning")
|
||||
dataset, dataloader = build_parquet_map_style_dataloader(
|
||||
args.path, args.batch_size, args.num_data_workers)
|
||||
|
||||
dataloader = build_parquet_map_style_dataloader(args.path,
|
||||
args.batch_size,
|
||||
args.num_data_workers)
|
||||
|
||||
for i, (latents, embeddings, masks,
|
||||
data_indices) in enumerate(dataloader):
|
||||
logger.info("Batch %d data_indices: %s", i, data_indices)
|
||||
# Second pass - verify latent sums match
|
||||
second_pass_sums = []
|
||||
for i, (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)",
|
||||
i, latent_sum, first_pass_sums[i])
|
||||
if i >= args.num_batches_per_epoch * 2 - 1:
|
||||
break
|
||||
|
||||
# Verify all sums match
|
||||
if all(
|
||||
abs(a - b) < 1e-6
|
||||
for a, b in zip(first_pass_sums, second_pass_sums)):
|
||||
logger.info(
|
||||
"All latent sums match between passes - resume verification successful!"
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
"Latent sums do not match between passes - resume verification failed!"
|
||||
)
|
||||
|
||||
start_time = time.time()
|
||||
total_samples = 0
|
||||
total_batches = 0
|
||||
for _ in range(args.num_epoch):
|
||||
for i, (latents, embeddings, masks,
|
||||
data_indices) in enumerate(dataloader):
|
||||
caption_text) in enumerate(dataloader):
|
||||
if i >= args.num_batches_per_epoch:
|
||||
break
|
||||
|
||||
|
||||
@@ -26,15 +26,47 @@ pyarrow_schema_i2v = pa.schema([
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
pa.field("text_attention_mask_bytes", pa.binary()),
|
||||
# e.g., [SeqLen]
|
||||
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bool' or 'int8'
|
||||
pa.field("text_attention_mask_dtype", pa.string()),
|
||||
#I2V
|
||||
pa.field("clip_feature_bytes", pa.binary()),
|
||||
pa.field("clip_feature_shape", pa.list_(pa.int64())),
|
||||
pa.field("clip_feature_dtype", pa.string()),
|
||||
pa.field("first_frame_latent_bytes", pa.binary()),
|
||||
pa.field("first_frame_latent_shape", pa.list_(pa.int64())),
|
||||
pa.field("first_frame_latent_dtype", pa.string()),
|
||||
# I2V Validation
|
||||
pa.field("pil_image_bytes", pa.binary()),
|
||||
pa.field("pil_image_shape", pa.list_(pa.int64())),
|
||||
pa.field("pil_image_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
pa.field("media_type", pa.string()), # 'image' or 'video'
|
||||
pa.field("width", pa.int64()),
|
||||
pa.field("height", pa.int64()),
|
||||
# -- Video-specific (can be null/default for images) ---
|
||||
# Number of frames processed (e.g., 1 for image, N for video)
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
pyarrow_schema_i2v_validation = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
#I2V
|
||||
pa.field("clip_feature_bytes", pa.binary()),
|
||||
pa.field("clip_feature_shape", pa.list_(pa.int64())),
|
||||
pa.field("clip_feature_dtype", pa.string()),
|
||||
# I2V Validation
|
||||
pa.field("pil_image_bytes", pa.binary()),
|
||||
pa.field("pil_image_shape", pa.list_(pa.int64())),
|
||||
pa.field("pil_image_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
@@ -64,11 +96,6 @@ pyarrow_schema_t2v = pa.schema([
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
pa.field("text_attention_mask_bytes", pa.binary()),
|
||||
# e.g., [SeqLen]
|
||||
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bool' or 'int8'
|
||||
pa.field("text_attention_mask_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
@@ -80,4 +107,26 @@ pyarrow_schema_t2v = pa.schema([
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
])
|
||||
|
||||
pyarrow_schema_t2v_validation = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
pa.field("media_type", pa.string()), # 'image' or 'video'
|
||||
pa.field("width", pa.int64()),
|
||||
pa.field("height", pa.int64()),
|
||||
# -- Video-specific (can be null/default for images) ---
|
||||
# Number of frames processed (e.g., 1 for image, N for video)
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
@@ -1,137 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from multiprocessing import Pool, cpu_count
|
||||
from pathlib import Path
|
||||
|
||||
import torchvision
|
||||
from tqdm import tqdm
|
||||
|
||||
|
||||
def get_video_info(video_path):
|
||||
"""Get video information using torchvision."""
|
||||
# Read video tensor (T, C, H, W)
|
||||
video_tensor, _, info = torchvision.io.read_video(str(video_path),
|
||||
output_format="TCHW",
|
||||
pts_unit="sec")
|
||||
|
||||
num_frames = video_tensor.shape[0]
|
||||
height = video_tensor.shape[2]
|
||||
width = video_tensor.shape[3]
|
||||
fps = info.get("video_fps", 0)
|
||||
duration = num_frames / fps if fps > 0 else 0
|
||||
|
||||
# Extract name
|
||||
_, _, videos_dir, video_name = str(video_path).split("/")
|
||||
|
||||
return {
|
||||
"path": str(video_name),
|
||||
"resolution": {
|
||||
"width": width,
|
||||
"height": height
|
||||
},
|
||||
"size": os.path.getsize(video_path),
|
||||
"fps": fps,
|
||||
"duration": duration,
|
||||
"num_frames": num_frames
|
||||
}
|
||||
|
||||
|
||||
def prepare_dataset_json(folder_path,
|
||||
output_name="videos2caption.json",
|
||||
num_workers=None) -> None:
|
||||
"""Prepare dataset information from a folder containing videos and prompt.txt."""
|
||||
folder_path = Path(folder_path)
|
||||
|
||||
# Read prompt file
|
||||
prompt_file = folder_path / "prompt.txt"
|
||||
if not prompt_file.exists():
|
||||
raise FileNotFoundError(f"prompt.txt not found in {folder_path}")
|
||||
|
||||
with open(prompt_file) as f:
|
||||
prompts = [line.strip() for line in f.readlines() if line.strip()]
|
||||
|
||||
# Read videos file
|
||||
videos_file = folder_path / "videos.txt"
|
||||
if not videos_file.exists():
|
||||
raise FileNotFoundError(f"videos.txt not found in {folder_path}")
|
||||
|
||||
with open(videos_file) as f:
|
||||
video_paths = [line.strip() for line in f.readlines() if line.strip()]
|
||||
|
||||
if len(prompts) != len(video_paths):
|
||||
raise ValueError(
|
||||
f"Number of prompts ({len(prompts)}) does not match number of videos ({len(video_paths)})"
|
||||
)
|
||||
|
||||
# Prepare arguments for multiprocessing
|
||||
process_args = [folder_path / video_path for video_path in video_paths]
|
||||
|
||||
# Determine number of workers
|
||||
if num_workers is None:
|
||||
num_workers = max(1, cpu_count() - 1) # Leave one CPU free
|
||||
|
||||
# Process videos in parallel
|
||||
start_time = time.time()
|
||||
with Pool(num_workers) as pool:
|
||||
results = list(
|
||||
tqdm(pool.imap(get_video_info, process_args),
|
||||
total=len(process_args),
|
||||
desc="Processing videos",
|
||||
unit="video"))
|
||||
|
||||
# Combine results with prompts
|
||||
dataset_info = []
|
||||
for result, prompt in zip(results, prompts):
|
||||
result["cap"] = [prompt]
|
||||
dataset_info.append(result)
|
||||
|
||||
# Calculate total processing time
|
||||
total_time = time.time() - start_time
|
||||
total_videos = len(dataset_info)
|
||||
avg_time_per_video = total_time / total_videos if total_videos > 0 else 0
|
||||
|
||||
print("\nProcessing completed:")
|
||||
print(f"Total videos processed: {total_videos}")
|
||||
print(f"Total time: {total_time:.2f} seconds")
|
||||
print(f"Average time per video: {avg_time_per_video:.2f} seconds")
|
||||
|
||||
# Save to JSON file
|
||||
output_file = folder_path / output_name
|
||||
with open(output_file, 'w') as f:
|
||||
json.dump(dataset_info, f, indent=2)
|
||||
|
||||
# Create merge.txt
|
||||
merge_file = folder_path / "merge.txt"
|
||||
with open(merge_file, 'w') as f:
|
||||
f.write(f"{folder_path}/videos,{output_file}\n")
|
||||
|
||||
print(f"Dataset information saved to {output_file}")
|
||||
print(f"Merge file created at {merge_file}")
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(
|
||||
description='Prepare video dataset information in JSON format')
|
||||
parser.add_argument(
|
||||
'--folder',
|
||||
type=str,
|
||||
required=True,
|
||||
help='Path to the folder containing videos and prompt.txt')
|
||||
parser.add_argument(
|
||||
'--output',
|
||||
type=str,
|
||||
default='videos2caption.json',
|
||||
help='Name of the output JSON file (default: videos2caption.json)')
|
||||
parser.add_argument('--workers',
|
||||
type=int,
|
||||
default=32,
|
||||
help='Number of worker processes (default: 16)')
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = parse_args()
|
||||
prepare_dataset_json(args.folder, args.output, args.workers)
|
||||
@@ -0,0 +1,278 @@
|
||||
import os
|
||||
import pickle
|
||||
import random
|
||||
from typing import Dict, List, Tuple
|
||||
|
||||
import numpy as np
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
import tqdm
|
||||
from torch.utils.data import IterableDataset, get_worker_info
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.v1.dataset.utils import collate_latents_embs_masks
|
||||
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
|
||||
get_world_size)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class BatchIterator:
|
||||
# TODO: Implement state_dict and load_state_dict to support resume.
|
||||
def __init__(self, files, batch_size, text_padding_length, keys,
|
||||
worker_num_samples, read_batch_size):
|
||||
self.files = files
|
||||
self.batch_size = batch_size
|
||||
self.text_padding_length = text_padding_length
|
||||
self.keys = keys
|
||||
self.worker_num_samples = worker_num_samples
|
||||
self.processed_samples = 0
|
||||
self.buffer = []
|
||||
self.read_batch_size = read_batch_size
|
||||
|
||||
def __iter__(self):
|
||||
for file in self.files:
|
||||
if self.processed_samples >= self.worker_num_samples:
|
||||
return
|
||||
|
||||
reader = pq.ParquetFile(file)
|
||||
for batch in reader.iter_batches(batch_size=self.read_batch_size):
|
||||
if self.processed_samples >= self.worker_num_samples:
|
||||
return
|
||||
|
||||
self.buffer.extend(batch.to_pylist())
|
||||
|
||||
while len(self.buffer) >= self.batch_size:
|
||||
if self.processed_samples >= self.worker_num_samples:
|
||||
return
|
||||
|
||||
batch_to_process = self.buffer[:self.batch_size]
|
||||
self.buffer = self.buffer[self.batch_size:]
|
||||
|
||||
all_latents, all_embs, all_masks, caption_text = collate_latents_embs_masks(
|
||||
batch_to_process, self.text_padding_length, self.keys)
|
||||
self.processed_samples += self.batch_size
|
||||
yield all_latents, all_embs, all_masks, caption_text
|
||||
|
||||
|
||||
class LatentsParquetIterStyleDataset(IterableDataset):
|
||||
"""Efficient loader for video-text data from a directory of Parquet files."""
|
||||
|
||||
# Modify this in the future if we want to add more keys, for example, in image to video.
|
||||
keys = [("vae_latent", "latent"), ("text_embedding")]
|
||||
|
||||
def __init__(self,
|
||||
path: str,
|
||||
batch_size: int = 1024,
|
||||
cfg_rate: float = 0.1,
|
||||
num_workers: int = 1,
|
||||
drop_last: bool = True,
|
||||
text_padding_length: int = 512,
|
||||
seed: int = 42,
|
||||
read_batch_size: int = 32,
|
||||
parquet_schema: pa.Schema = None):
|
||||
super().__init__()
|
||||
self.path = str(path)
|
||||
self.batch_size = batch_size
|
||||
self.parquet_schema = parquet_schema
|
||||
self.cfg_rate = cfg_rate
|
||||
self.text_padding_length = text_padding_length
|
||||
self.seed = seed
|
||||
self.read_batch_size = read_batch_size
|
||||
# Get distributed training info
|
||||
self.global_rank = get_world_rank()
|
||||
self.world_size = get_world_size()
|
||||
self.sp_world_size = get_sp_world_size()
|
||||
self.num_sp_groups = self.world_size // self.sp_world_size
|
||||
num_workers = 1 if num_workers == 0 else num_workers
|
||||
# Get sharding info
|
||||
shard_parquet_files, shard_total_samples, shard_parquet_lengths = shard_parquet_files_across_sp_groups_and_workers(
|
||||
self.path, self.num_sp_groups, num_workers, seed)
|
||||
|
||||
if drop_last:
|
||||
self.worker_num_samples = min(
|
||||
shard_total_samples) // batch_size * batch_size
|
||||
# Assign files to current rank's SP group
|
||||
ith_sp_group = self.global_rank // self.sp_world_size
|
||||
self.sp_group_parquet_files = shard_parquet_files[ith_sp_group::self
|
||||
.num_sp_groups]
|
||||
self.sp_group_parquet_lengths = shard_parquet_lengths[
|
||||
ith_sp_group::self.num_sp_groups]
|
||||
self.sp_group_num_samples = shard_total_samples[ith_sp_group::self.
|
||||
num_sp_groups]
|
||||
logger.info(
|
||||
"In total %d parquet files, %d samples, after sharding we retain %d samples due to drop_last",
|
||||
sum([len(shard) for shard in shard_parquet_files]),
|
||||
sum(shard_total_samples),
|
||||
self.worker_num_samples * self.num_sp_groups * num_workers)
|
||||
else:
|
||||
raise ValueError("drop_last must be True")
|
||||
logger.info("Each dataloader worker will load %d samples",
|
||||
self.worker_num_samples)
|
||||
|
||||
def __iter__(self):
|
||||
worker_info = get_worker_info()
|
||||
worker_id = worker_info.id if worker_info is not None else 1
|
||||
|
||||
worker_files = self.sp_group_parquet_files[worker_id]
|
||||
|
||||
batch_iterator = BatchIterator(
|
||||
files=worker_files,
|
||||
batch_size=self.batch_size,
|
||||
text_padding_length=self.text_padding_length,
|
||||
keys=self.keys,
|
||||
worker_num_samples=self.worker_num_samples,
|
||||
read_batch_size=self.read_batch_size) # type: ignore
|
||||
|
||||
yield from batch_iterator
|
||||
|
||||
if batch_iterator.processed_samples != self.worker_num_samples:
|
||||
raise ValueError(
|
||||
"Rank %d, Worker %d: Not enough samples to process, this should not happen",
|
||||
self.global_rank, worker_id)
|
||||
|
||||
|
||||
def shard_parquet_files_across_sp_groups_and_workers(
|
||||
path: str,
|
||||
num_sp_groups: int,
|
||||
num_workers: int,
|
||||
seed: int = 42,
|
||||
) -> Tuple[List[List[str]], List[int], List[Dict[str, int]]]:
|
||||
"""
|
||||
Shard parquet files across SP groups and workers in a balanced way.
|
||||
|
||||
Args:
|
||||
path: Directory containing parquet files
|
||||
num_sp_groups: Number of SP groups to shard across
|
||||
num_workers: Number of workers per SP group
|
||||
seed: Random seed for shuffling
|
||||
|
||||
Returns:
|
||||
Tuple containing:
|
||||
- List of lists of parquet files for each shard
|
||||
- List of total samples per shard
|
||||
- List of dictionaries mapping file paths to their lengths
|
||||
"""
|
||||
# Check if sharding plan already exists
|
||||
sharding_info_dir = os.path.join(
|
||||
path, f"sharding_info_{num_sp_groups}_sp_groups_{num_workers}_workers")
|
||||
if os.path.exists(sharding_info_dir):
|
||||
logger.info("Sharding plan already exists")
|
||||
logger.info("Loading sharding plan from %s", sharding_info_dir)
|
||||
try:
|
||||
with open(
|
||||
os.path.join(sharding_info_dir, "shard_parquet_files.pkl"),
|
||||
"rb") as f:
|
||||
shard_parquet_files = pickle.load(f)
|
||||
with open(
|
||||
os.path.join(sharding_info_dir, "shard_total_samples.pkl"),
|
||||
"rb") as f:
|
||||
shard_total_samples = pickle.load(f)
|
||||
with open(
|
||||
os.path.join(sharding_info_dir,
|
||||
"shard_parquet_lengths.pkl"), "rb") as f:
|
||||
shard_parquet_lengths = pickle.load(f)
|
||||
return shard_parquet_files, shard_total_samples, shard_parquet_lengths
|
||||
except Exception as e:
|
||||
logger.error("Error loading sharding plan: %s", str(e))
|
||||
logger.info("Falling back to creating new sharding plan")
|
||||
|
||||
if get_world_rank() == 0:
|
||||
logger.info("Scanning for parquet files in %s", path)
|
||||
|
||||
# Find all parquet files
|
||||
parquet_files = []
|
||||
|
||||
for root, _, files in os.walk(path):
|
||||
for file in files:
|
||||
if file.endswith('.parquet'):
|
||||
parquet_files.append(os.path.join(root, file))
|
||||
|
||||
if not parquet_files:
|
||||
raise ValueError("No parquet files found in %s", path)
|
||||
|
||||
# Calculate file lengths efficiently using a single pass
|
||||
logger.info("Calculating file lengths...")
|
||||
lengths = []
|
||||
for file in tqdm.tqdm(parquet_files, desc="Reading parquet files"):
|
||||
lengths.append(pq.ParquetFile(file).metadata.num_rows)
|
||||
|
||||
total_samples = sum(lengths)
|
||||
logger.info("Found %d files with %d total samples", len(parquet_files),
|
||||
total_samples)
|
||||
|
||||
# Sort files by length for better balancing
|
||||
sorted_indices = np.argsort(lengths)
|
||||
sorted_files = [parquet_files[i] for i in sorted_indices]
|
||||
sorted_lengths = [lengths[i] for i in sorted_indices]
|
||||
|
||||
# Create shards
|
||||
num_shards = num_sp_groups * num_workers
|
||||
shard_parquet_files = [[] for _ in range(num_shards)]
|
||||
shard_total_samples = [0] * num_shards
|
||||
shard_parquet_lengths = [{} for _ in range(num_shards)]
|
||||
|
||||
# Distribute files to shards using a greedy approach
|
||||
logger.info("Distributing files to shards...")
|
||||
for file, length in zip(reversed(sorted_files),
|
||||
reversed(sorted_lengths)):
|
||||
# Find shard with minimum current length
|
||||
target_shard = np.argmin(shard_total_samples)
|
||||
shard_parquet_files[target_shard].append(file)
|
||||
shard_total_samples[target_shard] += length
|
||||
shard_parquet_lengths[target_shard][file] = length
|
||||
#randomize each shard
|
||||
for shard in shard_parquet_files:
|
||||
random.seed(seed)
|
||||
random.shuffle(shard)
|
||||
|
||||
save_dir = os.path.join(
|
||||
path,
|
||||
f"sharding_info_{num_sp_groups}_sp_groups_{num_workers}_workers")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
with open(os.path.join(save_dir, "shard_parquet_files.pkl"), "wb") as f:
|
||||
pickle.dump(shard_parquet_files, f)
|
||||
with open(os.path.join(save_dir, "shard_total_samples.pkl"), "wb") as f:
|
||||
pickle.dump(shard_total_samples, f)
|
||||
with open(os.path.join(save_dir, "shard_parquet_lengths.pkl"),
|
||||
"wb") as f:
|
||||
pickle.dump(shard_parquet_lengths, f)
|
||||
logger.info("Saved sharding info to %s", save_dir)
|
||||
|
||||
# wait for all ranks to finish
|
||||
torch.distributed.barrier()
|
||||
# recursive call
|
||||
return shard_parquet_files_across_sp_groups_and_workers(
|
||||
path, num_sp_groups, num_workers, seed)
|
||||
|
||||
|
||||
def build_parquet_iterable_style_dataloader(
|
||||
path: str,
|
||||
batch_size: int,
|
||||
num_data_workers: int,
|
||||
cfg_rate: float = 0.0,
|
||||
drop_last: bool = True,
|
||||
text_padding_length: int = 512,
|
||||
seed: int = 42,
|
||||
read_batch_size: int = 32
|
||||
) -> Tuple[LatentsParquetIterStyleDataset, StatefulDataLoader]:
|
||||
"""Build a dataloader for the LatentsParquetIterStyleDataset."""
|
||||
dataset = LatentsParquetIterStyleDataset(
|
||||
path=path,
|
||||
batch_size=batch_size,
|
||||
cfg_rate=cfg_rate,
|
||||
num_workers=num_data_workers,
|
||||
drop_last=drop_last,
|
||||
text_padding_length=text_padding_length,
|
||||
seed=seed,
|
||||
read_batch_size=read_batch_size)
|
||||
|
||||
loader = StatefulDataLoader(
|
||||
dataset,
|
||||
batch_size=1,
|
||||
num_workers=num_data_workers,
|
||||
pin_memory=True,
|
||||
)
|
||||
return dataset, loader
|
||||
@@ -1,15 +1,18 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import pickle
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
import numpy as np
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
# Torch in general
|
||||
import torch
|
||||
import tqdm
|
||||
# Dataset
|
||||
from torch.utils.data import Dataset, Sampler
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.v1.dataset.utils import collate_rows_from_parquet_schema
|
||||
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
|
||||
get_world_size)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -30,6 +33,7 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
|
||||
sp_world_size: int,
|
||||
global_rank: int,
|
||||
drop_last: bool = True,
|
||||
drop_first_row: bool = False,
|
||||
seed: int = 0,
|
||||
):
|
||||
self.batch_size = batch_size
|
||||
@@ -45,6 +49,11 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
|
||||
# Create a random permutation of all indices
|
||||
global_indices = torch.randperm(self.dataset_size, generator=rng)
|
||||
|
||||
if drop_first_row:
|
||||
# drop 0 in global_indices
|
||||
global_indices = global_indices[global_indices != 0]
|
||||
self.dataset_size = self.dataset_size - 1
|
||||
|
||||
if self.drop_last:
|
||||
# For drop_last=True, we:
|
||||
# 1. Ensure total samples is divisible by (batch_size * num_sp_groups)
|
||||
@@ -56,19 +65,22 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
|
||||
self.num_sp_groups *
|
||||
self.batch_size]
|
||||
else:
|
||||
# add more indices to make it divisible by (batch_size * num_sp_groups)
|
||||
padding_size = self.num_sp_groups * self.batch_size - (
|
||||
self.dataset_size % (self.num_sp_groups * self.batch_size))
|
||||
global_indices = torch.cat(
|
||||
[global_indices, global_indices[:padding_size]])
|
||||
if self.dataset_size % (self.num_sp_groups * self.batch_size) != 0:
|
||||
# add more indices to make it divisible by (batch_size * num_sp_groups)
|
||||
padding_size = self.num_sp_groups * self.batch_size - (
|
||||
self.dataset_size % (self.num_sp_groups * self.batch_size))
|
||||
logger.info("Padding the dataset from %d to %d",
|
||||
self.dataset_size, self.dataset_size + padding_size)
|
||||
global_indices = torch.cat(
|
||||
[global_indices, global_indices[:padding_size]])
|
||||
|
||||
# shard the indices to each sp group
|
||||
ith_sp_group = self.global_rank // self.sp_world_size
|
||||
sp_group_local_indices = global_indices[ith_sp_group::self.
|
||||
num_sp_groups]
|
||||
|
||||
self.sp_group_local_indices = sp_group_local_indices
|
||||
logger.info("sp_group_local_indices: %d", len(sp_group_local_indices))
|
||||
logger.info("Dataset size for each sp group: %d",
|
||||
len(sp_group_local_indices))
|
||||
|
||||
def __iter__(self):
|
||||
indices = self.sp_group_local_indices
|
||||
@@ -81,20 +93,49 @@ class DP_SP_BatchSampler(Sampler[List[int]]):
|
||||
|
||||
|
||||
def get_parquet_files_and_length(path: str):
|
||||
lengths = []
|
||||
file_names = []
|
||||
for root, _, files in os.walk(path):
|
||||
for file in sorted(files):
|
||||
if file.endswith('.parquet'):
|
||||
file_path = os.path.join(root, file)
|
||||
num_rows = pq.ParquetFile(file_path).metadata.num_rows
|
||||
lengths.append(num_rows)
|
||||
file_names.append(file_path)
|
||||
# sort according to file name to ensure all rank has the same order (in case os.walk is not sorted)
|
||||
file_names_sorted, lengths_sorted = zip(
|
||||
*sorted(zip(file_names, lengths), key=lambda x: x[0]))
|
||||
assert len(file_names_sorted) != 0, "No parquet files found in the dataset"
|
||||
return file_names_sorted, lengths_sorted
|
||||
# Check if cached info exists
|
||||
cache_dir = os.path.join(path, "map_style_cache")
|
||||
cache_file = os.path.join(cache_dir, "file_info.pkl")
|
||||
|
||||
if os.path.exists(cache_file):
|
||||
logger.info("Loading cached file info from %s", cache_file)
|
||||
try:
|
||||
with open(cache_file, "rb") as f:
|
||||
file_names_sorted, lengths_sorted = pickle.load(f)
|
||||
return file_names_sorted, lengths_sorted
|
||||
except Exception as e:
|
||||
logger.error("Error loading cached file info: %s", str(e))
|
||||
logger.info("Falling back to scanning files")
|
||||
|
||||
# If no cache exists or loading failed, scan files
|
||||
if get_world_rank() == 0:
|
||||
lengths = []
|
||||
file_names = []
|
||||
for root, _, files in os.walk(path):
|
||||
for file in sorted(files):
|
||||
if file.endswith('.parquet'):
|
||||
file_path = os.path.join(root, file)
|
||||
file_names.append(file_path)
|
||||
for file_path in tqdm.tqdm(file_names,
|
||||
desc="Reading parquet files to get lengths"):
|
||||
num_rows = pq.ParquetFile(file_path).metadata.num_rows
|
||||
lengths.append(num_rows)
|
||||
# sort according to file name to ensure all rank has the same order (in case os.walk is not sorted)
|
||||
file_names_sorted, lengths_sorted = zip(
|
||||
*sorted(zip(file_names, lengths), key=lambda x: x[0]))
|
||||
assert len(
|
||||
file_names_sorted) != 0, "No parquet files found in the dataset"
|
||||
|
||||
os.makedirs(cache_dir, exist_ok=True)
|
||||
with open(cache_file, "wb") as f:
|
||||
pickle.dump((file_names_sorted, lengths_sorted), f)
|
||||
logger.info("Saved file info to %s", cache_file)
|
||||
|
||||
# Wait for rank 0 to finish saving
|
||||
if get_world_size() > 1:
|
||||
torch.distributed.barrier()
|
||||
|
||||
return get_parquet_files_and_length(path)
|
||||
|
||||
|
||||
def read_row_from_parquet_file(parquet_files: List[str], global_row_idx: int,
|
||||
@@ -145,20 +186,24 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
Using parquet for map style dataset is not efficient, we mainly keep it for backward compatibility and debugging.
|
||||
"""
|
||||
# Modify this in the future if we want to add more keys, for example, in image to video.
|
||||
keys = ["vae_latent", "text_embedding"]
|
||||
keys = [("vae_latent", "latent"), "text_embedding", "clip_feature",
|
||||
"first_frame_latent", "pil_image"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path: str,
|
||||
batch_size: int,
|
||||
parquet_schema: pa.Schema,
|
||||
cfg_rate: float = 0.0,
|
||||
seed: int = 42,
|
||||
drop_last: bool = True,
|
||||
drop_first_row: bool = False,
|
||||
text_padding_length: int = 512,
|
||||
):
|
||||
super().__init__()
|
||||
self.path = path
|
||||
self.cfg_rate = cfg_rate
|
||||
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"
|
||||
@@ -168,15 +213,6 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
self.parquet_files, self.lengths = get_parquet_files_and_length(path)
|
||||
self.batch = batch_size
|
||||
self.text_padding_length = text_padding_length
|
||||
self._cols = [
|
||||
"vae_latent_bytes",
|
||||
"vae_latent_shape",
|
||||
"text_embedding_bytes",
|
||||
"text_embedding_shape",
|
||||
"text_embedding_dtype",
|
||||
"height",
|
||||
"width",
|
||||
]
|
||||
self.sampler = DP_SP_BatchSampler(
|
||||
batch_size=batch_size,
|
||||
dataset_size=sum(self.lengths),
|
||||
@@ -184,27 +220,14 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
sp_world_size=get_sp_world_size(),
|
||||
global_rank=get_world_rank(),
|
||||
drop_last=drop_last,
|
||||
drop_first_row=drop_first_row,
|
||||
seed=seed,
|
||||
)
|
||||
logger.info("Dataset initialized with %d parquet files and %d rows",
|
||||
len(self.parquet_files), sum(self.lengths))
|
||||
|
||||
def _get_torch_tensors_from_row_dict(
|
||||
self, row_dict: Dict[str, Any]) -> Dict[str, torch.Tensor]:
|
||||
"""
|
||||
Get the latents and prompts from a row dictionary.
|
||||
"""
|
||||
return_dict = {}
|
||||
for key in self.keys:
|
||||
shape = row_dict[f"{key}_shape"]
|
||||
bytes = row_dict[f"{key}_bytes"]
|
||||
# TODO (peiyuan): read precision
|
||||
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
|
||||
data = torch.from_numpy(data)
|
||||
return_dict[key] = data
|
||||
return return_dict
|
||||
|
||||
def get_validation_negative_prompt(self) -> tuple[Any, Any, Any, Any]:
|
||||
def get_validation_negative_prompt(
|
||||
self) -> tuple[torch.Tensor, torch.Tensor, str]:
|
||||
"""
|
||||
Get the negative prompt for validation.
|
||||
This method ensures the negative prompt is loaded and cached properly.
|
||||
@@ -218,39 +241,22 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
row_dict = read_row_from_parquet_file([file_path], row_idx,
|
||||
[self.lengths[0]])
|
||||
|
||||
# Get tensors using the existing helper method
|
||||
data = self._get_torch_tensors_from_row_dict(row_dict)
|
||||
emb = data["text_embedding"]
|
||||
batch = collate_rows_from_parquet_schema([row_dict],
|
||||
self.parquet_schema,
|
||||
self.text_padding_length)
|
||||
negative_prompt = batch['info_list'][0]['prompt']
|
||||
negative_prompt_embedding = batch['text_embedding']
|
||||
negative_prompt_attention_mask = batch['text_attention_mask']
|
||||
if len(negative_prompt_embedding.shape) == 2:
|
||||
negative_prompt_embedding = negative_prompt_embedding.unsqueeze(0)
|
||||
if len(negative_prompt_attention_mask.shape) == 1:
|
||||
negative_prompt_attention_mask = negative_prompt_attention_mask.unsqueeze(
|
||||
0).unsqueeze(0)
|
||||
|
||||
# Pad the embedding and get mask
|
||||
padded_emb, mask = self._pad(emb, self.text_padding_length)
|
||||
|
||||
# Pin memory for faster transfer to GPU
|
||||
padded_emb = padded_emb
|
||||
mask = mask
|
||||
|
||||
return None, padded_emb, mask, None
|
||||
|
||||
def _pad(self, t: torch.Tensor, padding_length: int) -> torch.Tensor:
|
||||
"""
|
||||
Pad or crop an embedding [L, D] to exactly padding_length tokens.
|
||||
Return:
|
||||
- [L, D] tensor in pinned CPU memory
|
||||
- [L] attention mask in pinned CPU memory
|
||||
"""
|
||||
L, D = t.shape
|
||||
if padding_length > L: # pad
|
||||
pad = torch.zeros(padding_length - L,
|
||||
D,
|
||||
dtype=t.dtype,
|
||||
device=t.device)
|
||||
return torch.cat([t, pad], 0), torch.cat(
|
||||
[torch.ones(L), torch.zeros(padding_length - L)], 0)
|
||||
else: # crop
|
||||
return t[:padding_length], torch.ones(padding_length)
|
||||
return negative_prompt_embedding, negative_prompt_attention_mask, negative_prompt
|
||||
|
||||
# PyTorch calls this ONLY because the batch_sampler yields a list
|
||||
def __getitems__(self, indices: List[int]):
|
||||
def __getitems__(self, indices: List[int]) -> Dict[str, Any]:
|
||||
"""
|
||||
Batch fetch using read_row_from_parquet_file for each index.
|
||||
"""
|
||||
@@ -259,29 +265,12 @@ class LatentsParquetMapStyleDataset(Dataset):
|
||||
for idx in indices
|
||||
]
|
||||
|
||||
# Initialize tensors to hold padded embeddings and masks
|
||||
all_latents = []
|
||||
all_embs = []
|
||||
all_masks = []
|
||||
|
||||
# Process each row individually
|
||||
for i, row in enumerate(rows):
|
||||
# Get tensors from row
|
||||
data = self._get_torch_tensors_from_row_dict(row)
|
||||
latents, emb = data["vae_latent"], data["text_embedding"]
|
||||
|
||||
padded_emb, mask = self._pad(emb, self.text_padding_length)
|
||||
# Store in batch tensors
|
||||
all_latents.append(latents)
|
||||
all_embs.append(padded_emb)
|
||||
all_masks.append(mask)
|
||||
|
||||
# Pin memory for faster transfer to GPU
|
||||
all_latents = torch.stack(all_latents)
|
||||
all_embs = torch.stack(all_embs)
|
||||
all_masks = torch.stack(all_masks)
|
||||
|
||||
return all_latents, all_embs, all_masks, indices
|
||||
# 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):
|
||||
return sum(self.lengths)
|
||||
@@ -298,8 +287,10 @@ def build_parquet_map_style_dataloader(
|
||||
path,
|
||||
batch_size,
|
||||
num_data_workers,
|
||||
parquet_schema,
|
||||
cfg_rate=0.0,
|
||||
drop_last=True,
|
||||
drop_first_row=False,
|
||||
text_padding_length=512,
|
||||
seed=42) -> Tuple[LatentsParquetMapStyleDataset, StatefulDataLoader]:
|
||||
dataset = LatentsParquetMapStyleDataset(
|
||||
@@ -307,7 +298,9 @@ def build_parquet_map_style_dataloader(
|
||||
batch_size,
|
||||
cfg_rate=cfg_rate,
|
||||
drop_last=drop_last,
|
||||
drop_first_row=drop_first_row,
|
||||
text_padding_length=text_padding_length,
|
||||
parquet_schema=parquet_schema,
|
||||
seed=seed)
|
||||
|
||||
loader = StatefulDataLoader(
|
||||
|
||||
@@ -0,0 +1,592 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import Counter
|
||||
from dataclasses import dataclass
|
||||
from os.path import join as opj
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class PreprocessBatch:
|
||||
"""
|
||||
Batch information for dataset processing stages.
|
||||
|
||||
This class holds all the information about a video-caption or image-caption pair
|
||||
as it moves through the processing pipeline. Fields are populated by different stages.
|
||||
"""
|
||||
# Raw metadata
|
||||
path: str
|
||||
cap: Union[str, List[str]]
|
||||
resolution: Optional[Dict] = None
|
||||
fps: Optional[float] = None
|
||||
duration: Optional[float] = None
|
||||
|
||||
# Processed metadata
|
||||
num_frames: Optional[int] = None
|
||||
sample_frame_index: Optional[List[int]] = None
|
||||
sample_num_frames: Optional[int] = None
|
||||
|
||||
# Processed data
|
||||
pixel_values: Optional[torch.Tensor] = None
|
||||
text: Optional[str] = None
|
||||
input_ids: Optional[torch.Tensor] = None
|
||||
cond_mask: Optional[torch.Tensor] = None
|
||||
|
||||
@property
|
||||
def is_video(self) -> bool:
|
||||
"""Check if this is a video item."""
|
||||
return self.path.endswith(".mp4")
|
||||
|
||||
@property
|
||||
def is_image(self) -> bool:
|
||||
"""Check if this is an image item."""
|
||||
return self.path.endswith(".jpg")
|
||||
|
||||
|
||||
class DatasetStage(ABC):
|
||||
"""
|
||||
Abstract base class for dataset processing stages.
|
||||
|
||||
Similar to PipelineStage but designed for dataset preprocessing operations.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Process the dataset batch.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch to process
|
||||
**kwargs: Additional processing parameters
|
||||
|
||||
Returns:
|
||||
Processed batch
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class DatasetFilterStage(ABC):
|
||||
"""
|
||||
Abstract base class for dataset filtering stages.
|
||||
|
||||
These stages can filter out items during metadata processing.
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
|
||||
"""
|
||||
Check if batch should be kept.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch to check
|
||||
**kwargs: Additional parameters
|
||||
|
||||
Returns:
|
||||
True if batch should be kept, False otherwise
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Process the dataset batch (for non-filtering operations).
|
||||
|
||||
Args:
|
||||
batch: Dataset batch to process
|
||||
**kwargs: Additional processing parameters
|
||||
|
||||
Returns:
|
||||
Processed batch
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class DataValidationStage(DatasetFilterStage):
|
||||
"""Stage for validating data items."""
|
||||
|
||||
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
|
||||
"""
|
||||
Validate data item.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch to validate
|
||||
|
||||
Returns:
|
||||
True if valid, False if invalid
|
||||
"""
|
||||
# Check for caption
|
||||
if batch.cap is None:
|
||||
return False
|
||||
|
||||
if batch.is_video:
|
||||
# Validate video-specific fields
|
||||
if batch.duration is None or batch.fps is None:
|
||||
return False
|
||||
elif not batch.is_image:
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""Process does nothing for validation - filtering is handled by should_keep."""
|
||||
return batch
|
||||
|
||||
|
||||
class ResolutionFilterStage(DatasetFilterStage):
|
||||
"""Stage for filtering data items based on resolution constraints."""
|
||||
|
||||
def __init__(self,
|
||||
max_h_div_w_ratio: float = 17 / 16,
|
||||
min_h_div_w_ratio: float = 8 / 16,
|
||||
max_height: int = 1024,
|
||||
max_width: int = 1024):
|
||||
self.max_h_div_w_ratio = max_h_div_w_ratio
|
||||
self.min_h_div_w_ratio = min_h_div_w_ratio
|
||||
self.max_height = max_height
|
||||
self.max_width = max_width
|
||||
|
||||
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
|
||||
"""
|
||||
Check if data item passes resolution filtering.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch with resolution information
|
||||
|
||||
Returns:
|
||||
True if passes filter, False otherwise
|
||||
"""
|
||||
# Only apply to videos
|
||||
if not batch.is_video:
|
||||
return True
|
||||
|
||||
if batch.resolution is None:
|
||||
return False
|
||||
|
||||
height = batch.resolution.get("height", None)
|
||||
width = batch.resolution.get("width", None)
|
||||
if height is None or width is None:
|
||||
return False
|
||||
|
||||
# Check aspect ratio
|
||||
aspect = self.max_height / self.max_width
|
||||
hw_aspect_thr = 1.5
|
||||
|
||||
return self.filter_resolution(
|
||||
height,
|
||||
width,
|
||||
max_h_div_w_ratio=hw_aspect_thr * aspect,
|
||||
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
|
||||
)
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""Process does nothing for resolution filtering - filtering is handled by should_keep."""
|
||||
return batch
|
||||
|
||||
def filter_resolution(self, h: int, w: int, max_h_div_w_ratio: float,
|
||||
min_h_div_w_ratio: float) -> bool:
|
||||
"""Filter based on height/width ratio."""
|
||||
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
|
||||
|
||||
|
||||
class FrameSamplingStage(DatasetFilterStage):
|
||||
"""Stage for temporal frame sampling and indexing."""
|
||||
|
||||
def __init__(self,
|
||||
num_frames: int,
|
||||
train_fps: int,
|
||||
speed_factor: int = 1,
|
||||
video_length_tolerance_range: float = 5.0,
|
||||
drop_short_ratio: float = 0.0):
|
||||
self.num_frames = num_frames
|
||||
self.train_fps = train_fps
|
||||
self.speed_factor = speed_factor
|
||||
self.video_length_tolerance_range = video_length_tolerance_range
|
||||
self.drop_short_ratio = drop_short_ratio
|
||||
|
||||
def should_keep(self, batch: PreprocessBatch, **kwargs) -> bool:
|
||||
"""
|
||||
Check if video should be kept based on length constraints.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch
|
||||
|
||||
Returns:
|
||||
True if should be kept, False otherwise
|
||||
"""
|
||||
if batch.is_image:
|
||||
return True
|
||||
|
||||
if batch.duration is None or batch.fps is None:
|
||||
return False
|
||||
|
||||
num_frames = math.ceil(batch.fps * batch.duration)
|
||||
|
||||
# Check if video is too long
|
||||
if (num_frames / batch.fps > self.video_length_tolerance_range *
|
||||
(self.num_frames / self.train_fps * self.speed_factor)):
|
||||
return False
|
||||
|
||||
# Resample frame indices to check length
|
||||
frame_interval = batch.fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, num_frames,
|
||||
frame_interval).astype(int)
|
||||
|
||||
# Filter short videos
|
||||
return not (len(frame_indices) < self.num_frames
|
||||
and random.random() < self.drop_short_ratio)
|
||||
|
||||
def process(self,
|
||||
batch: PreprocessBatch,
|
||||
temporal_sample_fn=None,
|
||||
**kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Process frame sampling for video data items.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch
|
||||
temporal_sample_fn: Function for temporal sampling
|
||||
|
||||
Returns:
|
||||
Updated batch with frame sampling info
|
||||
"""
|
||||
if batch.is_image:
|
||||
# For images, just add sample info
|
||||
batch.sample_frame_index = [0]
|
||||
batch.sample_num_frames = 1
|
||||
return batch
|
||||
|
||||
assert batch.duration is not None and batch.fps is not None
|
||||
batch.num_frames = math.ceil(batch.fps * batch.duration)
|
||||
|
||||
# Resample frame indices
|
||||
frame_interval = batch.fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, batch.num_frames,
|
||||
frame_interval).astype(int)
|
||||
|
||||
# Temporal crop if too long
|
||||
if len(frame_indices
|
||||
) > self.num_frames and temporal_sample_fn is not None:
|
||||
begin_index, end_index = temporal_sample_fn(len(frame_indices))
|
||||
frame_indices = frame_indices[begin_index:end_index]
|
||||
|
||||
batch.sample_frame_index = frame_indices.tolist()
|
||||
batch.sample_num_frames = len(frame_indices)
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
class VideoTransformStage(DatasetStage):
|
||||
"""Stage for video data transformation."""
|
||||
|
||||
def __init__(self, transform) -> None:
|
||||
self.transform = transform
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Transform video data.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch with video information
|
||||
|
||||
Returns:
|
||||
Batch with transformed video tensor
|
||||
"""
|
||||
if not batch.is_video:
|
||||
return batch
|
||||
|
||||
assert os.path.exists(batch.path), f"file {batch.path} do not exist!"
|
||||
assert batch.sample_frame_index is not None, "Frame indices must be set before transformation"
|
||||
|
||||
torchvision_video, _, metadata = torchvision.io.read_video(
|
||||
batch.path, output_format="TCHW")
|
||||
video = torchvision_video[batch.sample_frame_index]
|
||||
if self.transform is not None:
|
||||
video = self.transform(video)
|
||||
video = rearrange(video, "t c h w -> c t h w")
|
||||
video = video.to(torch.uint8)
|
||||
|
||||
h, w = video.shape[-2:]
|
||||
assert (
|
||||
h / w <= 17 / 16 and h / w >= 8 / 16
|
||||
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({batch.path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
|
||||
|
||||
video = video.float() / 127.5 - 1.0
|
||||
batch.pixel_values = video
|
||||
return batch
|
||||
|
||||
|
||||
class ImageTransformStage(DatasetStage):
|
||||
"""Stage for image data transformation."""
|
||||
|
||||
def __init__(self, transform, transform_topcrop) -> None:
|
||||
self.transform = transform
|
||||
self.transform_topcrop = transform_topcrop
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Transform image data.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch with image information
|
||||
|
||||
Returns:
|
||||
Batch with transformed image tensor
|
||||
"""
|
||||
if not batch.is_image:
|
||||
return batch
|
||||
|
||||
image = Image.open(batch.path).convert("RGB")
|
||||
image = torch.from_numpy(np.array(image))
|
||||
image = rearrange(image, "h w c -> c h w").unsqueeze(0)
|
||||
|
||||
if self.transform_topcrop is not None:
|
||||
image = self.transform_topcrop(image)
|
||||
elif self.transform is not None:
|
||||
image = self.transform(image)
|
||||
|
||||
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
|
||||
image = image.float() / 127.5 - 1.0
|
||||
batch.pixel_values = image
|
||||
return batch
|
||||
|
||||
|
||||
class TextEncodingStage(DatasetStage):
|
||||
"""Stage for text tokenization and encoding."""
|
||||
|
||||
def __init__(self, tokenizer, text_max_length: int, cfg_rate: float = 0.0):
|
||||
self.tokenizer = tokenizer
|
||||
self.text_max_length = text_max_length
|
||||
self.cfg_rate = cfg_rate
|
||||
|
||||
def process(self, batch: PreprocessBatch, **kwargs) -> PreprocessBatch:
|
||||
"""
|
||||
Process text data.
|
||||
|
||||
Args:
|
||||
batch: Dataset batch with caption information
|
||||
|
||||
Returns:
|
||||
Batch with encoded text information
|
||||
"""
|
||||
text = batch.cap
|
||||
if not isinstance(text, list):
|
||||
text = [text]
|
||||
text = [random.choice(text)]
|
||||
|
||||
text = text[0] if random.random() > self.cfg_rate else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
|
||||
batch.text = text
|
||||
batch.input_ids = text_tokens_and_mask["input_ids"]
|
||||
batch.cond_mask = text_tokens_and_mask["attention_mask"]
|
||||
return batch
|
||||
|
||||
|
||||
class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
|
||||
torch.distributed.checkpoint.stateful.Stateful):
|
||||
"""
|
||||
Merged dataset for video and caption data with stage-based processing.
|
||||
|
||||
This dataset processes video and image data through a series of stages:
|
||||
- Data validation
|
||||
- Resolution filtering
|
||||
- Frame sampling
|
||||
- Transformation
|
||||
- Text encoding
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
data_merge_path: str,
|
||||
args,
|
||||
transform,
|
||||
temporal_sample,
|
||||
transform_topcrop,
|
||||
start_idx: int = 0):
|
||||
self.data_merge_path = data_merge_path
|
||||
self.start_idx = start_idx
|
||||
self.args = args
|
||||
self.temporal_sample = temporal_sample
|
||||
|
||||
# Initialize tokenizer
|
||||
tokenizer_path = os.path.join(args.model_path, "tokenizer")
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
|
||||
cache_dir=args.cache_dir)
|
||||
|
||||
# Initialize processing stages
|
||||
self._init_stages(args, transform, transform_topcrop, tokenizer)
|
||||
|
||||
# Process metadata
|
||||
self.processed_batches = self._process_metadata()
|
||||
|
||||
def _init_stages(self, args, transform, transform_topcrop,
|
||||
tokenizer) -> None:
|
||||
"""Initialize all processing stages."""
|
||||
self.validation_stage = DataValidationStage()
|
||||
self.resolution_filter_stage = ResolutionFilterStage(
|
||||
max_height=args.max_height, max_width=args.max_width)
|
||||
self.frame_sampling_stage = FrameSamplingStage(
|
||||
num_frames=args.num_frames,
|
||||
train_fps=args.train_fps,
|
||||
speed_factor=args.speed_factor,
|
||||
video_length_tolerance_range=args.video_length_tolerance_range,
|
||||
drop_short_ratio=args.drop_short_ratio)
|
||||
self.video_transform_stage = VideoTransformStage(transform)
|
||||
self.image_transform_stage = ImageTransformStage(
|
||||
transform, transform_topcrop)
|
||||
self.text_encoding_stage = TextEncodingStage(
|
||||
tokenizer=tokenizer,
|
||||
text_max_length=args.text_max_length,
|
||||
cfg_rate=args.cfg)
|
||||
|
||||
def _load_raw_data(self) -> List[Dict]:
|
||||
"""Load raw data from JSON files."""
|
||||
all_data = []
|
||||
|
||||
# Read folder-annotation pairs
|
||||
with open(self.data_merge_path) as f:
|
||||
folder_anno_pairs = [
|
||||
line.strip().split(",") for line in f if line.strip()
|
||||
]
|
||||
|
||||
# Process each folder-annotation pair
|
||||
for folder, annotation_file in folder_anno_pairs:
|
||||
with open(annotation_file) as f:
|
||||
data_items = json.load(f)
|
||||
|
||||
# Update paths with folder prefix
|
||||
for item in data_items:
|
||||
item["path"] = opj(folder, item["path"])
|
||||
|
||||
all_data.extend(data_items)
|
||||
|
||||
return all_data[self.start_idx:]
|
||||
|
||||
def _process_metadata(self) -> List[PreprocessBatch]:
|
||||
"""Process the raw metadata through all filtering stages."""
|
||||
raw_data = self._load_raw_data()
|
||||
processed_batches = []
|
||||
|
||||
# Initialize counters
|
||||
filter_counts = {
|
||||
"validation_failed": 0,
|
||||
"resolution_failed": 0,
|
||||
"frame_sampling_failed": 0
|
||||
}
|
||||
sample_num_frames: List[int] = []
|
||||
|
||||
for item in raw_data:
|
||||
batch = PreprocessBatch(path=item["path"],
|
||||
cap=item["cap"],
|
||||
resolution=item.get("resolution"),
|
||||
fps=item.get("fps"),
|
||||
duration=item.get("duration"))
|
||||
|
||||
# Apply filtering stages
|
||||
if not self._apply_filter_stages(batch, filter_counts):
|
||||
continue
|
||||
|
||||
# Apply frame sampling processing
|
||||
batch = self.frame_sampling_stage.process(
|
||||
batch, temporal_sample_fn=self.temporal_sample)
|
||||
|
||||
processed_batches.append(batch)
|
||||
assert batch.sample_num_frames is not None
|
||||
sample_num_frames.append(batch.sample_num_frames)
|
||||
|
||||
self._log_filtering_stats(filter_counts, sample_num_frames,
|
||||
len(raw_data), len(processed_batches))
|
||||
return processed_batches
|
||||
|
||||
def _apply_filter_stages(self, batch: PreprocessBatch,
|
||||
filter_counts: Dict[str, int]) -> bool:
|
||||
"""Apply all filter stages and update counters. Returns True if batch should be kept."""
|
||||
if not self.validation_stage.should_keep(batch):
|
||||
filter_counts["validation_failed"] += 1
|
||||
return False
|
||||
|
||||
if not self.resolution_filter_stage.should_keep(batch):
|
||||
filter_counts["resolution_failed"] += 1
|
||||
return False
|
||||
|
||||
if not self.frame_sampling_stage.should_keep(batch):
|
||||
filter_counts["frame_sampling_failed"] += 1
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def _log_filtering_stats(self, filter_counts: Dict[str, int],
|
||||
sample_num_frames: List[int], before_count: int,
|
||||
after_count: int):
|
||||
"""Log filtering statistics."""
|
||||
logger.info(
|
||||
"validation_failed: %d, resolution_failed: %d, frame_sampling_failed: %d, "
|
||||
"Counter(sample_num_frames): %s, before filter: %d, after filter: %d",
|
||||
filter_counts['validation_failed'],
|
||||
filter_counts['resolution_failed'],
|
||||
filter_counts['frame_sampling_failed'], Counter(sample_num_frames),
|
||||
before_count, after_count)
|
||||
|
||||
def __iter__(self):
|
||||
"""Iterate through processed data items."""
|
||||
for idx in range(len(self.processed_batches)):
|
||||
yield self._get_item(idx)
|
||||
|
||||
def __len__(self):
|
||||
return len(self.processed_batches)
|
||||
|
||||
def _get_item(self, idx: int) -> Dict:
|
||||
"""Get a single processed data item."""
|
||||
batch = self.processed_batches[idx]
|
||||
|
||||
# Apply transformation stages
|
||||
batch = self.video_transform_stage.process(batch)
|
||||
batch = self.image_transform_stage.process(batch)
|
||||
batch = self.text_encoding_stage.process(batch)
|
||||
|
||||
# Build result dictionary
|
||||
result = {
|
||||
"pixel_values": batch.pixel_values,
|
||||
"text": batch.text,
|
||||
"input_ids": batch.input_ids,
|
||||
"cond_mask": batch.cond_mask,
|
||||
"path": batch.path,
|
||||
}
|
||||
|
||||
# Add video-specific fields
|
||||
if batch.is_video:
|
||||
result.update({"fps": batch.fps, "duration": batch.duration})
|
||||
|
||||
return result
|
||||
|
||||
def state_dict(self) -> Dict[str, Any]:
|
||||
"""Return state dict for checkpointing."""
|
||||
return {"processed_batches": self.processed_batches}
|
||||
|
||||
def load_state_dict(self, state_dict: Dict[str, Any]) -> None:
|
||||
"""Load state dict from checkpoint."""
|
||||
self.processed_batches = state_dict["processed_batches"]
|
||||
@@ -1,352 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
import random
|
||||
from collections import Counter
|
||||
from os.path import join as opj
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
from PIL import Image
|
||||
from torch.utils.data import Dataset
|
||||
|
||||
from fastvideo.utils.dataset_utils import DecordInit
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
|
||||
|
||||
class SingletonMeta(type):
|
||||
_instances: dict[type, 'SingletonMeta'] = {}
|
||||
|
||||
def __call__(cls, *args, **kwargs):
|
||||
if cls not in cls._instances:
|
||||
instance = super().__call__(*args, **kwargs)
|
||||
cls._instances[cls] = instance
|
||||
return cls._instances[cls]
|
||||
|
||||
|
||||
class DataSetProg(metaclass=SingletonMeta):
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.cap_list: list[dict] = []
|
||||
self.elements: list[int] = []
|
||||
self.num_workers = 1
|
||||
self.n_elements = 0
|
||||
self.worker_elements: dict[int, list[int]] = {}
|
||||
self.n_used_elements: dict[int, int] = {}
|
||||
|
||||
def set_cap_list(self, num_workers, cap_list, n_elements) -> None:
|
||||
self.num_workers = num_workers
|
||||
self.cap_list = cap_list
|
||||
self.n_elements = n_elements
|
||||
self.elements = list(range(n_elements))
|
||||
random.shuffle(self.elements)
|
||||
print(f"n_elements: {len(self.elements)}", flush=True)
|
||||
|
||||
for i in range(self.num_workers):
|
||||
self.n_used_elements[i] = 0
|
||||
per_worker = int(
|
||||
math.ceil(len(self.elements) / float(self.num_workers)))
|
||||
start = i * per_worker
|
||||
end = min(start + per_worker, len(self.elements))
|
||||
self.worker_elements[i] = self.elements[start:end]
|
||||
|
||||
def get_item(self, work_info) -> int:
|
||||
worker_id = 0 if work_info is None else work_info.id
|
||||
|
||||
idx = self.worker_elements[worker_id][
|
||||
self.n_used_elements[worker_id] %
|
||||
len(self.worker_elements[worker_id])]
|
||||
self.n_used_elements[worker_id] += 1
|
||||
return idx
|
||||
|
||||
|
||||
dataset_prog = DataSetProg()
|
||||
|
||||
|
||||
def filter_resolution(h: int,
|
||||
w: int,
|
||||
max_h_div_w_ratio: float = 17 / 16,
|
||||
min_h_div_w_ratio: float = 8 / 16) -> bool:
|
||||
return h / w <= max_h_div_w_ratio and h / w >= min_h_div_w_ratio
|
||||
|
||||
|
||||
class T2V_dataset(Dataset):
|
||||
|
||||
def __init__(self,
|
||||
args,
|
||||
transform,
|
||||
temporal_sample,
|
||||
tokenizer,
|
||||
transform_topcrop,
|
||||
start_idx=0) -> None:
|
||||
self.start_idx = start_idx
|
||||
self.data = args.data_merge_path
|
||||
self.num_frames = args.num_frames
|
||||
self.train_fps = args.train_fps
|
||||
self.use_image_num = args.use_image_num
|
||||
self.transform = transform
|
||||
self.transform_topcrop = transform_topcrop
|
||||
self.temporal_sample = temporal_sample
|
||||
self.tokenizer = tokenizer
|
||||
self.text_max_length = args.text_max_length
|
||||
self.cfg = args.cfg
|
||||
self.speed_factor = args.speed_factor
|
||||
self.max_height = args.max_height
|
||||
self.max_width = args.max_width
|
||||
self.drop_short_ratio = args.drop_short_ratio
|
||||
assert self.speed_factor >= 1
|
||||
self.v_decoder = DecordInit()
|
||||
self.video_length_tolerance_range = args.video_length_tolerance_range
|
||||
self.support_Chinese = True
|
||||
if "mt5" not in args.text_encoder_name:
|
||||
self.support_Chinese = False
|
||||
|
||||
cap_list = self.get_cap_list()
|
||||
|
||||
assert len(cap_list) > 0
|
||||
cap_list, self.sample_num_frames = self.define_frame_index(cap_list)
|
||||
self.lengths = self.sample_num_frames
|
||||
|
||||
n_elements = len(cap_list)
|
||||
dataset_prog.set_cap_list(args.dataloader_num_workers, cap_list,
|
||||
n_elements)
|
||||
|
||||
print(f"video length: {len(dataset_prog.cap_list)}", flush=True)
|
||||
|
||||
def set_checkpoint(self, n_used_elements):
|
||||
for i in range(len(dataset_prog.n_used_elements)):
|
||||
dataset_prog.n_used_elements[i] = n_used_elements
|
||||
|
||||
def __len__(self):
|
||||
return dataset_prog.n_elements
|
||||
|
||||
def __getitem__(self, idx):
|
||||
|
||||
data = self.get_data(idx)
|
||||
return data
|
||||
|
||||
def get_data(self, idx) -> dict:
|
||||
path = dataset_prog.cap_list[idx]["path"]
|
||||
if path.endswith(".mp4"):
|
||||
return self.get_video(idx)
|
||||
else:
|
||||
return self.get_image(idx)
|
||||
|
||||
def get_video(self, idx) -> dict:
|
||||
video_path = dataset_prog.cap_list[idx]["path"]
|
||||
assert os.path.exists(video_path), f"file {video_path} do not exist!"
|
||||
frame_indices = dataset_prog.cap_list[idx]["sample_frame_index"]
|
||||
|
||||
torchvision_video, _, metadata = torchvision.io.read_video(
|
||||
video_path, output_format="TCHW")
|
||||
video = torchvision_video[frame_indices]
|
||||
video = self.transform(video)
|
||||
video = rearrange(video, "t c h w -> c t h w")
|
||||
video = video.to(torch.uint8)
|
||||
assert video.dtype == torch.uint8
|
||||
|
||||
h, w = video.shape[-2:]
|
||||
assert (
|
||||
h / w <= 17 / 16 and h / w >= 8 / 16
|
||||
), f"Only videos with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But video ({video_path}) found ratio is {round(h / w, 2)} with the shape of {video.shape}"
|
||||
|
||||
video = video.float() / 127.5 - 1.0
|
||||
|
||||
text = dataset_prog.cap_list[idx]["cap"]
|
||||
if not isinstance(text, list):
|
||||
text = [text]
|
||||
text = [random.choice(text)]
|
||||
|
||||
text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"]
|
||||
cond_mask = text_tokens_and_mask["attention_mask"]
|
||||
return dict(pixel_values=video,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=video_path,
|
||||
fps=dataset_prog.cap_list[idx]["fps"],
|
||||
duration=dataset_prog.cap_list[idx]["duration"])
|
||||
|
||||
def get_image(self, idx) -> dict:
|
||||
image_data = dataset_prog.cap_list[
|
||||
idx] # [{'path': path, 'cap': cap}, ...]
|
||||
|
||||
image = Image.open(image_data["path"]).convert("RGB") # [h, w, c]
|
||||
image = torch.from_numpy(np.array(image)) # [h, w, c]
|
||||
image = rearrange(image, "h w c -> c h w").unsqueeze(0) # [1 c h w]
|
||||
# for i in image:
|
||||
# h, w = i.shape[-2:]
|
||||
# assert h / w <= 17 / 16 and h / w >= 8 / 16, f'Only image with a ratio (h/w) less than 17/16 and more than 8/16 are supported. But found ratio is {round(h / w, 2)} with the shape of {i.shape}'
|
||||
|
||||
image = (self.transform_topcrop(image) if "human_images"
|
||||
in image_data["path"] else self.transform(image)
|
||||
) # [1 C H W] -> num_img [1 C H W]
|
||||
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
|
||||
|
||||
image = image.float() / 127.5 - 1.0
|
||||
|
||||
caps: list[str] = (image_data["cap"] if isinstance(
|
||||
image_data["cap"], list) else [image_data["cap"]])
|
||||
caps = [random.choice(caps)]
|
||||
text = caps
|
||||
input_ids, cond_mask = [], []
|
||||
single_text = text[0] if random.random() > self.cfg else ""
|
||||
text_tokens_and_mask = self.tokenizer(
|
||||
single_text,
|
||||
max_length=self.text_max_length,
|
||||
padding="max_length",
|
||||
truncation=True,
|
||||
return_attention_mask=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
input_ids = text_tokens_and_mask["input_ids"] # 1, l
|
||||
cond_mask = text_tokens_and_mask["attention_mask"] # 1, l
|
||||
return dict(
|
||||
pixel_values=image,
|
||||
text=text,
|
||||
input_ids=input_ids,
|
||||
cond_mask=cond_mask,
|
||||
path=image_data["path"],
|
||||
)
|
||||
|
||||
def define_frame_index(self, cap_list) -> tuple[list[dict], list[int]]:
|
||||
new_cap_list = []
|
||||
sample_num_frames = []
|
||||
cnt_too_long = 0
|
||||
cnt_too_short = 0
|
||||
cnt_no_cap = 0
|
||||
cnt_no_resolution = 0
|
||||
cnt_resolution_mismatch = 0
|
||||
cnt_movie = 0
|
||||
cnt_img = 0
|
||||
for i in cap_list:
|
||||
path = i["path"]
|
||||
cap = i.get("cap", None)
|
||||
# ======no caption=====
|
||||
if cap is None:
|
||||
cnt_no_cap += 1
|
||||
continue
|
||||
if path.endswith(".mp4"):
|
||||
# ======no fps and duration=====
|
||||
duration = i.get("duration", None)
|
||||
fps = i.get("fps", None)
|
||||
if fps is None or duration is None:
|
||||
continue
|
||||
|
||||
# ======resolution mismatch=====
|
||||
resolution = i.get("resolution", None)
|
||||
if resolution is None:
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
else:
|
||||
if (resolution.get("height", None) is None
|
||||
or resolution.get("width", None) is None):
|
||||
cnt_no_resolution += 1
|
||||
continue
|
||||
height, width = i["resolution"]["height"], i["resolution"][
|
||||
"width"]
|
||||
aspect = self.max_height / self.max_width
|
||||
hw_aspect_thr = 1.5
|
||||
is_pick = filter_resolution(
|
||||
height,
|
||||
width,
|
||||
max_h_div_w_ratio=hw_aspect_thr * aspect,
|
||||
min_h_div_w_ratio=1 / hw_aspect_thr * aspect,
|
||||
)
|
||||
if not is_pick:
|
||||
print("resolution mismatch")
|
||||
cnt_resolution_mismatch += 1
|
||||
continue
|
||||
|
||||
# if path == 'finetrainers/3dgs-dissolve/videos/1.mp4':
|
||||
# from IPython import embed; embed()
|
||||
i["num_frames"] = math.ceil(fps * duration)
|
||||
# max 5.0 and min 1.0 are just thresholds to filter some videos which have suitable duration.
|
||||
if i["num_frames"] / fps > self.video_length_tolerance_range * (
|
||||
self.num_frames / self.train_fps * self.speed_factor
|
||||
): # too long video is not suitable for this training stage (self.num_frames)
|
||||
cnt_too_long += 1
|
||||
continue
|
||||
|
||||
# resample in case high fps, such as 50/60/90/144 -> train_fps(e.g, 24)
|
||||
frame_interval = fps / self.train_fps
|
||||
start_frame_idx = 0
|
||||
frame_indices = np.arange(start_frame_idx, i["num_frames"],
|
||||
frame_interval).astype(int)
|
||||
|
||||
# comment out it to enable dynamic frames training
|
||||
if (len(frame_indices) < self.num_frames
|
||||
and random.random() < self.drop_short_ratio):
|
||||
cnt_too_short += 1
|
||||
continue
|
||||
|
||||
# too long video will be temporal-crop randomly
|
||||
if len(frame_indices) > self.num_frames:
|
||||
begin_index, end_index = self.temporal_sample(
|
||||
len(frame_indices))
|
||||
frame_indices = frame_indices[begin_index:end_index]
|
||||
# frame_indices = frame_indices[:self.num_frames] # head crop
|
||||
i["sample_frame_index"] = frame_indices.tolist()
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = len(
|
||||
i["sample_frame_index"]
|
||||
) # will use in dataloader(group sampler)
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
elif path.endswith(".jpg"): # image
|
||||
cnt_img += 1
|
||||
new_cap_list.append(i)
|
||||
i["sample_num_frames"] = 1
|
||||
sample_num_frames.append(i["sample_num_frames"])
|
||||
else:
|
||||
raise NameError(
|
||||
f"Unknown file extension {path.split('.')[-1]}, only support .mp4 for video and .jpg for image"
|
||||
)
|
||||
# import ipdb;ipdb.set_trace()
|
||||
main_print(
|
||||
f"no_cap: {cnt_no_cap}, too_long: {cnt_too_long}, too_short: {cnt_too_short}, "
|
||||
f"no_resolution: {cnt_no_resolution}, resolution_mismatch: {cnt_resolution_mismatch}, "
|
||||
f"Counter(sample_num_frames): {Counter(sample_num_frames)}, cnt_movie: {cnt_movie}, cnt_img: {cnt_img}, "
|
||||
f"before filter: {len(cap_list)}, after filter: {len(new_cap_list)}"
|
||||
)
|
||||
return new_cap_list, sample_num_frames
|
||||
|
||||
def decord_read(self, path, frame_indices) -> torch.Tensor:
|
||||
decord_vr = self.v_decoder(path)
|
||||
video_data = decord_vr.get_batch(frame_indices).asnumpy()
|
||||
video_data = torch.from_numpy(video_data)
|
||||
video_data = video_data.permute(0, 3, 1, 2) # (T, H, W, C) -> (T C H W)
|
||||
return video_data
|
||||
|
||||
def read_jsons(self, data) -> list[dict]:
|
||||
cap_lists = []
|
||||
with open(data) as f:
|
||||
folder_anno = [
|
||||
i.strip().split(",") for i in f.readlines()
|
||||
if len(i.strip()) > 0
|
||||
]
|
||||
print(folder_anno)
|
||||
for folder, anno in folder_anno:
|
||||
with open(anno) as f:
|
||||
sub_list = json.load(f)
|
||||
for i in range(len(sub_list)):
|
||||
sub_list[i]["path"] = opj(folder, sub_list[i]["path"])
|
||||
cap_lists += sub_list
|
||||
return cap_lists
|
||||
|
||||
def get_cap_list(self) -> list:
|
||||
cap_lists = self.read_jsons(self.data)[self.start_idx:]
|
||||
return cap_lists
|
||||
@@ -0,0 +1,247 @@
|
||||
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:
|
||||
"""
|
||||
Pad or crop an embedding [L, D] to exactly padding_length tokens.
|
||||
Return:
|
||||
- [L, D] tensor in pinned CPU memory
|
||||
- [L] attention mask in pinned CPU memory
|
||||
"""
|
||||
L, D = t.shape
|
||||
if padding_length > L: # pad
|
||||
pad = torch.zeros(padding_length - L, D, dtype=t.dtype, device=t.device)
|
||||
return torch.cat([t, pad], 0), torch.cat(
|
||||
[torch.ones(L), torch.zeros(padding_length - L)], 0)
|
||||
else: # crop
|
||||
return t[:padding_length], torch.ones(padding_length)
|
||||
|
||||
|
||||
def get_torch_tensors_from_row_dict(row_dict, keys) -> Dict[str, Any]:
|
||||
"""
|
||||
Get the latents and prompts from a row dictionary.
|
||||
"""
|
||||
return_dict = {}
|
||||
for key in keys:
|
||||
shape, bytes = None, None
|
||||
if isinstance(key, tuple):
|
||||
for k in key:
|
||||
try:
|
||||
shape = row_dict[f"{k}_shape"]
|
||||
bytes = row_dict[f"{k}_bytes"]
|
||||
except KeyError:
|
||||
continue
|
||||
key = key[0]
|
||||
if shape is None or bytes is None:
|
||||
raise ValueError(f"Key {key} not found in row_dict")
|
||||
else:
|
||||
try:
|
||||
shape = row_dict[f"{key}_shape"]
|
||||
bytes = row_dict[f"{key}_bytes"]
|
||||
except KeyError:
|
||||
continue
|
||||
|
||||
# TODO (peiyuan): read precision
|
||||
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
|
||||
return return_dict
|
||||
|
||||
|
||||
def collate_latents_embs_masks(
|
||||
batch_to_process, text_padding_length, keys
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, List[str], Dict[str, Any],
|
||||
List[Dict[str, Any]]]:
|
||||
# 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)
|
||||
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"])
|
||||
except KeyError:
|
||||
caption_text.append(row["caption"])
|
||||
|
||||
# Pin memory for faster transfer to GPU
|
||||
all_latents = torch.stack(all_latents)
|
||||
all_embs = torch.stack(all_embs)
|
||||
all_masks = torch.stack(all_masks)
|
||||
all_extra_latents = {
|
||||
"clip_feature": torch.stack(all_clip_features),
|
||||
"first_frame_latent": torch.stack(all_first_frame_latents),
|
||||
"pil_image": all_pil_images,
|
||||
}
|
||||
|
||||
return all_latents, all_embs, all_masks, caption_text, all_extra_latents, all_infos
|
||||
|
||||
|
||||
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.
|
||||
|
||||
Args:
|
||||
rows: List of row dictionaries from parquet files
|
||||
parquet_schema: PyArrow schema defining the structure of the data
|
||||
|
||||
Returns:
|
||||
Dict containing batched tensors and metadata
|
||||
"""
|
||||
if not rows:
|
||||
return {}
|
||||
|
||||
# Initialize containers for different data types
|
||||
batch_data = {}
|
||||
|
||||
# Get tensor and metadata field names from schema (fields ending with '_bytes')
|
||||
tensor_fields = []
|
||||
metadata_fields = []
|
||||
for field in parquet_schema.names:
|
||||
if field.endswith('_bytes'):
|
||||
shape_field = field.replace('_bytes', '_shape')
|
||||
dtype_field = field.replace('_bytes', '_dtype')
|
||||
tensor_name = field.replace('_bytes', '')
|
||||
tensor_fields.append(tensor_name)
|
||||
assert shape_field in parquet_schema.names, f"Shape field {shape_field} not found in schema for field {field}. Currently we only support *_bytes fields for tensors."
|
||||
assert dtype_field in parquet_schema.names, f"Dtype field {dtype_field} not found in schema for field {field}. Currently we only support *_bytes fields for tensors."
|
||||
elif not field.endswith('_shape') and not field.endswith('_dtype'):
|
||||
# Only add actual metadata fields, not the shape/dtype helper fields
|
||||
metadata_fields.append(field)
|
||||
|
||||
# Process each tensor field efficiently
|
||||
for tensor_name in tensor_fields:
|
||||
tensor_list = []
|
||||
|
||||
for row in rows:
|
||||
# Get tensor data from row using the existing helper function pattern
|
||||
shape_key = f"{tensor_name}_shape"
|
||||
bytes_key = f"{tensor_name}_bytes"
|
||||
|
||||
if shape_key in row and bytes_key in row:
|
||||
# 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]
|
||||
|
||||
if len(bytes_data) == 0:
|
||||
tensor = torch.zeros(0, dtype=torch.bfloat16)
|
||||
else:
|
||||
# Convert bytes to tensor using float32 as default
|
||||
# logger.info("len(bytes_data): %s", len(bytes_data))
|
||||
# logger.info("shape: %s", shape)
|
||||
|
||||
data = np.frombuffer(
|
||||
bytes_data, dtype=np.float32).reshape(shape).copy()
|
||||
tensor = torch.from_numpy(data)
|
||||
# if len(data.shape) == 3:
|
||||
# B, L, D = tensor.shape
|
||||
# assert B == 1, "Batch size must be 1"
|
||||
# tensor = tensor.squeeze(0)
|
||||
|
||||
tensor_list.append(tensor)
|
||||
else:
|
||||
# Handle missing tensor data
|
||||
tensor_list.append(torch.zeros(0, dtype=torch.bfloat16))
|
||||
|
||||
# Stack tensors with special handling for text embeddings
|
||||
if tensor_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))
|
||||
|
||||
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
|
||||
]
|
||||
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 efficiently into info_list
|
||||
info_list = []
|
||||
for row in rows:
|
||||
info = {}
|
||||
for field in metadata_fields:
|
||||
info[field] = row.get(field, "")
|
||||
|
||||
# Add prompt field for backward compatibility
|
||||
info["prompt"] = info.get("caption", "")
|
||||
info_list.append(info)
|
||||
|
||||
batch_data['info_list'] = info_list
|
||||
|
||||
# Add caption_text for backward compatibility
|
||||
if info_list and 'caption' in info_list[0]:
|
||||
batch_data['caption_text'] = [info['caption'] for info in info_list]
|
||||
|
||||
return batch_data
|
||||
@@ -0,0 +1,104 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# adapted from: https://github.com/a-r-r-o-w/finetrainers/blob/main/finetrainers/data/dataset.py
|
||||
import os
|
||||
import pathlib
|
||||
|
||||
import datasets
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vision_utils import load_image, load_video
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ValidationDataset(torch.utils.data.IterableDataset):
|
||||
|
||||
def __init__(self, filename: str):
|
||||
super().__init__()
|
||||
|
||||
self.filename = pathlib.Path(filename)
|
||||
# get directory of filename
|
||||
# TODO(will)
|
||||
self.dir = os.path.abspath(self.filename.parent)
|
||||
|
||||
if not self.filename.exists():
|
||||
raise FileNotFoundError(
|
||||
f"File {self.filename.as_posix()} does not exist")
|
||||
|
||||
if self.filename.suffix == ".csv":
|
||||
data = datasets.load_dataset("csv",
|
||||
data_files=self.filename.as_posix(),
|
||||
split="train")
|
||||
elif self.filename.suffix == ".json":
|
||||
data = datasets.load_dataset("json",
|
||||
data_files=self.filename.as_posix(),
|
||||
split="train",
|
||||
field="data")
|
||||
elif self.filename.suffix == ".parquet":
|
||||
data = datasets.load_dataset("parquet",
|
||||
data_files=self.filename.as_posix(),
|
||||
split="train")
|
||||
elif self.filename.suffix == ".arrow":
|
||||
data = datasets.load_dataset("arrow",
|
||||
data_files=self.filename.as_posix(),
|
||||
split="train")
|
||||
else:
|
||||
_SUPPORTED_FILE_FORMATS = [".csv", ".json", ".parquet", ".arrow"]
|
||||
raise ValueError(
|
||||
f"Unsupported file format {self.filename.suffix} for validation dataset. Supported formats are: {_SUPPORTED_FILE_FORMATS}"
|
||||
)
|
||||
|
||||
self._data = data.to_iterable_dataset()
|
||||
|
||||
def __iter__(self):
|
||||
for sample in self._data:
|
||||
# For consistency reasons, we mandate that "caption" is always present in the validation dataset.
|
||||
# However, since the model specifications use "prompt", we create an alias here.
|
||||
sample["prompt"] = sample["caption"]
|
||||
|
||||
# Load image or video if the path is provided
|
||||
# TODO(aryan): need to handle custom columns here for control conditions
|
||||
sample["image"] = None
|
||||
sample["video"] = None
|
||||
|
||||
if sample.get("image_path", None) is not None:
|
||||
image_path = sample["image_path"]
|
||||
image_path = os.path.join(self.dir, image_path)
|
||||
if not pathlib.Path(image_path).is_file(
|
||||
) and not image_path.startswith("http"):
|
||||
logger.warning("Image file %s does not exist.", image_path)
|
||||
else:
|
||||
sample["image"] = load_image(image_path)
|
||||
|
||||
if sample.get("video_path", None) is not None:
|
||||
video_path = sample["video_path"]
|
||||
video_path = os.path.join(self.dir, video_path)
|
||||
if not pathlib.Path(video_path).is_file(
|
||||
) and not video_path.startswith("http"):
|
||||
logger.warning("Video file %s does not exist.", video_path)
|
||||
else:
|
||||
sample["video"] = load_video(video_path)
|
||||
|
||||
if sample.get("control_image_path", None) is not None:
|
||||
control_image_path = sample["control_image_path"]
|
||||
control_image_path = os.path.join(self.dir, control_image_path)
|
||||
if not pathlib.Path(control_image_path).is_file(
|
||||
) and not control_image_path.startswith("http"):
|
||||
logger.warning("Control Image file %s does not exist.",
|
||||
control_image_path)
|
||||
else:
|
||||
sample["control_image"] = load_image(control_image_path)
|
||||
|
||||
if sample.get("control_video_path", None) is not None:
|
||||
control_video_path = sample["control_video_path"]
|
||||
control_video_path = os.path.join(self.dir, control_video_path)
|
||||
if not pathlib.Path(control_video_path).is_file(
|
||||
) and not control_video_path.startswith("http"):
|
||||
logger.warning("Control Video file %s does not exist.",
|
||||
control_video_path)
|
||||
else:
|
||||
sample["control_video"] = load_video(control_video_path)
|
||||
|
||||
sample = {k: v for k, v in sample.items() if v is not None}
|
||||
yield sample
|
||||
@@ -4,9 +4,9 @@
|
||||
import argparse
|
||||
import dataclasses
|
||||
import os
|
||||
from typing import Any, Dict, List, Optional, cast
|
||||
from typing import List, cast
|
||||
|
||||
from fastvideo import PipelineConfig, VideoGenerator
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
|
||||
from fastvideo.v1.entrypoints.cli.utils import RaiseNotImplementedAction
|
||||
@@ -37,8 +37,6 @@ class GenerateSubcommand(CLISubcommand):
|
||||
def cmd(self, args: argparse.Namespace) -> None:
|
||||
excluded_args = ['subparser', 'config', 'dispatch_function']
|
||||
|
||||
FastVideoArgs.from_cli_args(args)
|
||||
|
||||
provided_args = {}
|
||||
for k, v in vars(args).items():
|
||||
if (k not in excluded_args and v is not None
|
||||
@@ -66,27 +64,19 @@ class GenerateSubcommand(CLISubcommand):
|
||||
|
||||
init_args = {
|
||||
k: v
|
||||
for k, v in merged_args.items() if k in self.init_arg_names
|
||||
for k, v in merged_args.items()
|
||||
if k not in self.generation_arg_names
|
||||
}
|
||||
generation_args = {
|
||||
k: v
|
||||
for k, v in merged_args.items() if k in self.generation_arg_names
|
||||
}
|
||||
|
||||
pipeline_config = PipelineConfig.from_pretrained(
|
||||
merged_args['model_path'])
|
||||
|
||||
update_config_from_args(pipeline_config.dit_config, merged_args,
|
||||
"dit_config")
|
||||
update_config_from_args(pipeline_config.vae_config, merged_args,
|
||||
"vae_config")
|
||||
update_config_from_args(pipeline_config, merged_args)
|
||||
|
||||
model_path = init_args.pop('model_path')
|
||||
prompt = generation_args.pop('prompt')
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_path=model_path, **init_args, pipeline_config=pipeline_config)
|
||||
generator = VideoGenerator.from_pretrained(model_path=model_path,
|
||||
**init_args)
|
||||
|
||||
generator.generate_video(prompt=prompt, **generation_args)
|
||||
|
||||
@@ -132,34 +122,3 @@ class GenerateSubcommand(CLISubcommand):
|
||||
|
||||
def cmd_init() -> List[CLISubcommand]:
|
||||
return [GenerateSubcommand()]
|
||||
|
||||
|
||||
def update_config_from_args(config: Any,
|
||||
args_dict: Dict[str, Any],
|
||||
prefix: Optional[str] = None) -> None:
|
||||
"""
|
||||
Update configuration object from arguments dictionary.
|
||||
|
||||
Args:
|
||||
config: The configuration object to update
|
||||
args_dict: Dictionary containing arguments
|
||||
prefix: Prefix for the configuration parameters in the args_dict.
|
||||
If None, assumes direct attribute mapping without prefix.
|
||||
"""
|
||||
# Handle top-level attributes (no prefix)
|
||||
if prefix is None:
|
||||
for key, value in args_dict.items():
|
||||
if hasattr(config, key) and value is not None:
|
||||
if key == "text_encoder_precisions" and isinstance(value, list):
|
||||
setattr(config, key, tuple(value))
|
||||
else:
|
||||
setattr(config, key, value)
|
||||
return
|
||||
|
||||
# Handle nested attributes with prefix
|
||||
prefix_with_dot = f"{prefix}."
|
||||
for key, value in args_dict.items():
|
||||
if key.startswith(prefix_with_dot) and value is not None:
|
||||
attr_name = key[len(prefix_with_dot):]
|
||||
if hasattr(config, attr_name):
|
||||
setattr(config, attr_name, value)
|
||||
|
||||
@@ -18,8 +18,6 @@ import torch
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
|
||||
from fastvideo.v1.configs.pipelines import (PipelineConfig,
|
||||
get_pipeline_config_cls_for_name)
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -55,9 +53,6 @@ class VideoGenerator:
|
||||
model_path: str,
|
||||
device: Optional[str] = None,
|
||||
torch_dtype: Optional[torch.dtype] = None,
|
||||
pipeline_config: Optional[
|
||||
Union[str
|
||||
| PipelineConfig]] = None,
|
||||
**kwargs) -> "VideoGenerator":
|
||||
"""
|
||||
Create a video generator from a pretrained model.
|
||||
@@ -66,35 +61,17 @@ class VideoGenerator:
|
||||
model_path: Path or identifier for the pretrained model
|
||||
device: Device to load the model on (e.g., "cuda", "cuda:0", "cpu")
|
||||
torch_dtype: Data type for model weights (e.g., torch.float16)
|
||||
**kwargs: Additional arguments to customize model loading
|
||||
pipeline_config: Pipeline config to use for inference
|
||||
**kwargs: Additional arguments to customize model loading, set any FastVideoArgs or PipelineConfig attributes here.
|
||||
|
||||
Returns:
|
||||
The created video generator
|
||||
|
||||
Priority level: Default pipeline config < User's pipeline config < User's kwargs
|
||||
"""
|
||||
config = None
|
||||
# 1. If users provide a pipeline config, it will override the default pipeline config
|
||||
if isinstance(pipeline_config, PipelineConfig):
|
||||
config = pipeline_config
|
||||
else:
|
||||
config_cls = get_pipeline_config_cls_for_name(model_path)
|
||||
if config_cls is not None:
|
||||
config = config_cls()
|
||||
if isinstance(pipeline_config, str):
|
||||
config.load_from_json(pipeline_config)
|
||||
|
||||
# 2. If users also provide some kwargs, it will override the pipeline config.
|
||||
# The user kwargs shouldn't contain model config parameters!
|
||||
if config is None:
|
||||
logger.warning("No config found for model %s, using default config",
|
||||
model_path)
|
||||
config_args = kwargs
|
||||
else:
|
||||
config_args = shallow_asdict(config)
|
||||
config_args.update(kwargs)
|
||||
|
||||
fastvideo_args = FastVideoArgs(model_path=model_path, **config_args)
|
||||
# If users also provide some kwargs, it will override the FastVideoArgs and PipelineConfig.
|
||||
kwargs['model_path'] = model_path
|
||||
fastvideo_args = FastVideoArgs.from_kwargs(kwargs)
|
||||
|
||||
return cls.from_fastvideo_args(fastvideo_args)
|
||||
|
||||
@@ -150,16 +127,17 @@ class VideoGenerator:
|
||||
"""
|
||||
# Create a copy of inference args to avoid modifying the original
|
||||
fastvideo_args = self.fastvideo_args
|
||||
pipeline_config = fastvideo_args.pipeline_config
|
||||
|
||||
# Validate inputs
|
||||
if not isinstance(prompt, str):
|
||||
raise TypeError(
|
||||
f"`prompt` must be a string, but got {type(prompt)}")
|
||||
prompt = prompt.strip()
|
||||
|
||||
if sampling_param is None:
|
||||
sampling_param = SamplingParam.from_pretrained(
|
||||
fastvideo_args.model_path)
|
||||
|
||||
kwargs["prompt"] = prompt
|
||||
sampling_param.update(kwargs)
|
||||
|
||||
@@ -176,10 +154,10 @@ class VideoGenerator:
|
||||
f"height={sampling_param.height}, width={sampling_param.width}, "
|
||||
f"num_frames={sampling_param.num_frames}")
|
||||
|
||||
temporal_scale_factor = fastvideo_args.vae_config.arch_config.temporal_compression_ratio
|
||||
temporal_scale_factor = pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
num_frames = sampling_param.num_frames
|
||||
num_gpus = fastvideo_args.num_gpus
|
||||
use_temporal_scaling_frames = fastvideo_args.vae_config.use_temporal_scaling_frames
|
||||
use_temporal_scaling_frames = pipeline_config.vae_config.use_temporal_scaling_frames
|
||||
|
||||
# Adjust number of frames based on number of GPUs
|
||||
if use_temporal_scaling_frames:
|
||||
@@ -238,18 +216,18 @@ class VideoGenerator:
|
||||
num_videos_per_prompt: {sampling_param.num_videos_per_prompt}
|
||||
guidance_scale: {sampling_param.guidance_scale}
|
||||
n_tokens: {n_tokens}
|
||||
flow_shift: {fastvideo_args.flow_shift}
|
||||
embedded_guidance_scale: {fastvideo_args.embedded_cfg_scale}
|
||||
flow_shift: {fastvideo_args.pipeline_config.flow_shift}
|
||||
embedded_guidance_scale: {fastvideo_args.pipeline_config.embedded_cfg_scale}
|
||||
save_video: {sampling_param.save_video}
|
||||
output_path: {sampling_param.output_path}
|
||||
""" # type: ignore[attr-defined]
|
||||
logger.info(debug_str)
|
||||
|
||||
# Prepare batch
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
eta=0.0,
|
||||
n_tokens=n_tokens,
|
||||
VSA_sparsity=fastvideo_args.VSA_sparsity,
|
||||
extra={},
|
||||
)
|
||||
|
||||
|
||||
+92
-208
@@ -6,26 +6,32 @@ import argparse
|
||||
import dataclasses
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import field
|
||||
from typing import Any, Callable, List, Optional, Tuple
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig, STA_Mode
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import FlexibleArgumentParser, StoreBoolean
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def preprocess_text(prompt: str) -> str:
|
||||
return prompt
|
||||
def clean_cli_args(args: argparse.Namespace) -> Dict[str, Any]:
|
||||
"""
|
||||
Clean the arguments by removing the ones that not explicitly provided by the user.
|
||||
"""
|
||||
provided_args = {}
|
||||
for k, v in vars(args).items():
|
||||
if (v is not None and hasattr(args, '_provided')
|
||||
and k in args._provided):
|
||||
provided_args[k] = v
|
||||
|
||||
|
||||
def postprocess_text(output: Any) -> Any:
|
||||
raise NotImplementedError
|
||||
return provided_args
|
||||
|
||||
|
||||
# args for fastvideo framework
|
||||
@dataclasses.dataclass
|
||||
class FastVideoArgs:
|
||||
# Model and path configuration
|
||||
# Model and path configuration (for convenience)
|
||||
model_path: str
|
||||
|
||||
# Cache strategy
|
||||
@@ -48,66 +54,28 @@ class FastVideoArgs:
|
||||
hsdp_shard_dim: int = -1
|
||||
dist_timeout: Optional[int] = None # timeout for torch.distributed
|
||||
|
||||
# Video generation parameters
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: Optional[float] = None
|
||||
pipeline_config: PipelineConfig = field(default_factory=PipelineConfig)
|
||||
|
||||
output_type: str = "pil"
|
||||
|
||||
# DiT configuration
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
precision: str = "bf16"
|
||||
use_cpu_offload: bool = True
|
||||
use_fsdp_inference: bool = True
|
||||
|
||||
# VAE configuration
|
||||
vae_precision: str = "fp16"
|
||||
vae_tiling: bool = True # Might change in between forward passes
|
||||
vae_sp: bool = False # Might change in between forward passes
|
||||
# vae_scale_factor: Optional[int] = None # Deprecated
|
||||
vae_config: VAEConfig = field(default_factory=VAEConfig)
|
||||
|
||||
# Image encoder configuration
|
||||
image_encoder_precision: str = "fp32"
|
||||
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
|
||||
|
||||
# Text encoder configuration
|
||||
DEFAULT_TEXT_ENCODER_PRECISIONS = (
|
||||
"fp16",
|
||||
# "fp16",
|
||||
)
|
||||
text_encoder_precisions: Tuple[str, ...] = field(
|
||||
default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS)
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (EncoderConfig(), ))
|
||||
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda: (preprocess_text, ))
|
||||
postprocess_text_funcs: Tuple[Callable[[Any], Any], ...] = field(
|
||||
default_factory=lambda: (postprocess_text, ))
|
||||
|
||||
# STA parameters
|
||||
STA_mode: Optional[str] = None
|
||||
skip_time_steps: int = 15
|
||||
# LoRA parameters
|
||||
lora_path: Optional[str] = None
|
||||
lora_nickname: Optional[
|
||||
str] = "default" # for swapping adapters in the pipeline
|
||||
lora_target_names: Optional[List[
|
||||
str]] = None # can restrict list of layers to adapt, e.g. ["q_proj"]
|
||||
|
||||
# STA parameters
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
STA_mode: STA_Mode = STA_Mode.STA_INFERENCE
|
||||
skip_time_steps: int = 15
|
||||
|
||||
# Compilation
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
disable_autocast: bool = False
|
||||
|
||||
# StepVideo specific parameters
|
||||
pos_magic: Optional[str] = None
|
||||
neg_magic: Optional[str] = None
|
||||
timesteps_scale: Optional[bool] = None
|
||||
# VSA parameters
|
||||
VSA_sparsity: float = 0.0 # inference/validation sparsity
|
||||
|
||||
# Logging
|
||||
log_level: str = "info"
|
||||
# Stage verification
|
||||
enable_stage_verification: bool = True
|
||||
|
||||
@property
|
||||
def training_mode(self) -> bool:
|
||||
@@ -125,11 +93,6 @@ class FastVideoArgs:
|
||||
help=
|
||||
"The path of the model weights. This can be a local folder or a Hugging Face repo ID.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dit-weight",
|
||||
type=str,
|
||||
help="Path to the DiT model weights",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model-dir",
|
||||
type=str,
|
||||
@@ -175,14 +138,12 @@ class FastVideoArgs:
|
||||
help="The number of GPUs to use.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tensor-parallel-size",
|
||||
"--tp-size",
|
||||
type=int,
|
||||
default=FastVideoArgs.tp_size,
|
||||
help="The tensor parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sequence-parallel-size",
|
||||
"--sp-size",
|
||||
type=int,
|
||||
default=FastVideoArgs.sp_size,
|
||||
@@ -207,19 +168,7 @@ class FastVideoArgs:
|
||||
help="Set timeout for torch.distributed initialization.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--embedded-cfg-scale",
|
||||
type=float,
|
||||
default=FastVideoArgs.embedded_cfg_scale,
|
||||
help="Embedded CFG scale",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--flow-shift",
|
||||
"--shift",
|
||||
type=float,
|
||||
default=FastVideoArgs.flow_shift,
|
||||
help="Flow shift parameter",
|
||||
)
|
||||
# Output type
|
||||
parser.add_argument(
|
||||
"--output-type",
|
||||
type=str,
|
||||
@@ -228,62 +177,14 @@ class FastVideoArgs:
|
||||
help="Output type for the generated video",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--precision",
|
||||
type=str,
|
||||
default=FastVideoArgs.precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for the model",
|
||||
)
|
||||
|
||||
# VAE configuration
|
||||
parser.add_argument(
|
||||
"--vae-precision",
|
||||
type=str,
|
||||
default=FastVideoArgs.vae_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for VAE",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-tiling",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.vae_tiling,
|
||||
help="Enable VAE tiling",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-sp",
|
||||
action=StoreBoolean,
|
||||
help="Enable VAE spatial parallelism",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--text-encoder-precisions",
|
||||
nargs="+",
|
||||
type=str,
|
||||
default=FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for each text encoder",
|
||||
)
|
||||
|
||||
# Image encoder config
|
||||
parser.add_argument(
|
||||
"--image-encoder-precision",
|
||||
type=str,
|
||||
default=FastVideoArgs.image_encoder_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for image encoder",
|
||||
)
|
||||
|
||||
# STA parameters
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
parser.add_argument(
|
||||
"--STA-mode",
|
||||
type=str,
|
||||
default=FastVideoArgs.STA_mode,
|
||||
choices=[
|
||||
"STA_inference", "STA_searching", "STA_tuning",
|
||||
"STA_tuning_cfg", None
|
||||
],
|
||||
help="STA mode",
|
||||
default=FastVideoArgs.STA_mode.value,
|
||||
choices=[mode.value for mode in STA_Mode],
|
||||
help=
|
||||
"STA mode contains STA_inference, STA_searching, STA_tuning, STA_tuning_cfg, None",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--skip-time-steps",
|
||||
@@ -323,69 +224,50 @@ class FastVideoArgs:
|
||||
"Disable autocast for denoising loop and vae decoding in pipeline sampling",
|
||||
)
|
||||
|
||||
# VSA parameters
|
||||
parser.add_argument(
|
||||
"--pos_magic",
|
||||
type=str,
|
||||
default=FastVideoArgs.pos_magic,
|
||||
help="Positive magic prompt for sampling",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--neg_magic",
|
||||
type=str,
|
||||
default=FastVideoArgs.neg_magic,
|
||||
help="Negative magic prompt for sampling",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--timesteps_scale",
|
||||
type=bool,
|
||||
default=FastVideoArgs.timesteps_scale,
|
||||
help="Bool for applying scheduler scale in set_timesteps",
|
||||
"--VSA-sparsity",
|
||||
type=float,
|
||||
default=FastVideoArgs.VSA_sparsity,
|
||||
help="Validation sparsity for VSA",
|
||||
)
|
||||
|
||||
# Logging
|
||||
# Stage verification
|
||||
parser.add_argument(
|
||||
"--log-level",
|
||||
type=str,
|
||||
default=FastVideoArgs.log_level,
|
||||
help="The logging level of all loggers.",
|
||||
"--enable-stage-verification",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.enable_stage_verification,
|
||||
help="Enable input/output verification for pipeline stages",
|
||||
)
|
||||
|
||||
# Add VAE configuration arguments
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEConfig
|
||||
VAEConfig.add_cli_args(parser)
|
||||
|
||||
# Add DiT configuration arguments
|
||||
from fastvideo.v1.configs.models.dits.base import DiTConfig
|
||||
DiTConfig.add_cli_args(parser)
|
||||
# Add pipeline configuration arguments
|
||||
PipelineConfig.add_cli_args(parser)
|
||||
|
||||
return parser
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "FastVideoArgs":
|
||||
args.tp_size = args.tensor_parallel_size
|
||||
args.sp_size = args.sequence_parallel_size
|
||||
args.flow_shift = getattr(args, "shift", args.flow_shift)
|
||||
|
||||
provided_args = clean_cli_args(args)
|
||||
# Get all fields from the dataclass
|
||||
attrs = [attr.name for attr in dataclasses.fields(cls)]
|
||||
|
||||
# Create a dictionary of attribute values, with defaults for missing attributes
|
||||
kwargs = {}
|
||||
for attr in attrs:
|
||||
# Handle renamed attributes or those with multiple CLI names
|
||||
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
|
||||
kwargs[attr] = args.tensor_parallel_size
|
||||
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
|
||||
kwargs[attr] = args.sequence_parallel_size
|
||||
elif attr == 'flow_shift' and hasattr(args, 'shift'):
|
||||
kwargs[attr] = args.shift
|
||||
if attr == 'pipeline_config':
|
||||
pipeline_config = PipelineConfig.from_kwargs(provided_args)
|
||||
kwargs[attr] = pipeline_config
|
||||
# Use getattr with default value from the dataclass for potentially missing attributes
|
||||
else:
|
||||
default_value = getattr(cls, attr, None)
|
||||
value = getattr(args, attr, default_value)
|
||||
if value is not None:
|
||||
kwargs[attr] = value
|
||||
kwargs[attr] = value # type: ignore
|
||||
|
||||
return cls(**kwargs) # type: ignore
|
||||
|
||||
@classmethod
|
||||
def from_kwargs(cls, kwargs: Dict[str, Any]) -> "FastVideoArgs":
|
||||
kwargs['pipeline_config'] = PipelineConfig.from_kwargs(kwargs)
|
||||
return cls(**kwargs)
|
||||
|
||||
def check_fastvideo_args(self) -> None:
|
||||
@@ -414,33 +296,17 @@ class FastVideoArgs:
|
||||
f"tp_size ({self.tp_size}) must be equal to sp_size ({self.sp_size})"
|
||||
)
|
||||
|
||||
# Validate VAE spatial parallelism with VAE tiling
|
||||
if self.vae_sp and not self.vae_tiling:
|
||||
raise ValueError(
|
||||
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
|
||||
)
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.text_encoder_precisions):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
|
||||
)
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
)
|
||||
|
||||
if len(self.preprocess_text_funcs) != len(self.postprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
)
|
||||
|
||||
if self.enable_torch_compile and self.num_gpus > 1:
|
||||
logger.warning(
|
||||
"Currently torch compile does not work with multi-gpu. Setting enable_torch_compile to False"
|
||||
)
|
||||
self.enable_torch_compile = False
|
||||
|
||||
if self.pipeline_config is None:
|
||||
raise ValueError("pipeline_config is not set in FastVideoArgs")
|
||||
|
||||
self.pipeline_config.check_pipeline_config()
|
||||
|
||||
|
||||
_current_fastvideo_args = None
|
||||
|
||||
@@ -514,7 +380,6 @@ class TrainingArgs(FastVideoArgs):
|
||||
# text encoder & vae & diffusion model
|
||||
pretrained_model_name_or_path: str = ""
|
||||
dit_model_name_or_path: str = ""
|
||||
cache_dir: str = ""
|
||||
|
||||
# diffusion setting
|
||||
ema_decay: float = 0.0
|
||||
@@ -523,12 +388,14 @@ class TrainingArgs(FastVideoArgs):
|
||||
precondition_outputs: bool = False
|
||||
|
||||
# validation & logs
|
||||
validation_prompt_dir: str = ""
|
||||
validation_dataset_file: str = ""
|
||||
validation_preprocessed_path: str = ""
|
||||
validation_sampling_steps: str = ""
|
||||
validation_guidance_scale: str = ""
|
||||
validation_steps: float = 0.0
|
||||
log_validation: bool = False
|
||||
tracker_project_name: str = ""
|
||||
wandb_run_name: str = ""
|
||||
seed: Optional[int] = None
|
||||
|
||||
# output
|
||||
@@ -576,30 +443,29 @@ class TrainingArgs(FastVideoArgs):
|
||||
# master_weight_type
|
||||
master_weight_type: str = ""
|
||||
|
||||
# For fast checking in LoRA pipeline
|
||||
training_mode: bool = True
|
||||
# VSA training decay parameters
|
||||
VSA_decay_rate: float = 0.01 # decay rate -> 0.02
|
||||
VSA_decay_interval_steps: int = 1 # decay interval steps -> 50
|
||||
|
||||
@classmethod
|
||||
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
|
||||
provided_args = clean_cli_args(args)
|
||||
# Get all fields from the dataclass
|
||||
attrs = [attr.name for attr in dataclasses.fields(cls)]
|
||||
|
||||
logger.info(provided_args)
|
||||
# Create a dictionary of attribute values, with defaults for missing attributes
|
||||
kwargs = {}
|
||||
for attr in attrs:
|
||||
# Handle renamed attributes or those with multiple CLI names
|
||||
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
|
||||
kwargs[attr] = args.tensor_parallel_size
|
||||
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
|
||||
kwargs[attr] = args.sequence_parallel_size
|
||||
elif attr == 'flow_shift' and hasattr(args, 'shift'):
|
||||
kwargs[attr] = args.shift
|
||||
if attr == 'pipeline_config':
|
||||
pipeline_config = PipelineConfig.from_kwargs(provided_args)
|
||||
kwargs[attr] = pipeline_config
|
||||
# Use getattr with default value from the dataclass for potentially missing attributes
|
||||
else:
|
||||
default_value = getattr(cls, attr, None)
|
||||
if getattr(args, attr, default_value) is not None:
|
||||
kwargs[attr] = getattr(args, attr, default_value)
|
||||
value = getattr(args, attr, default_value)
|
||||
kwargs[attr] = value # type: ignore
|
||||
|
||||
return cls(**kwargs)
|
||||
return cls(**kwargs) # type: ignore
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
@@ -671,9 +537,12 @@ class TrainingArgs(FastVideoArgs):
|
||||
help="Whether to precondition the outputs of the model")
|
||||
|
||||
# Validation and logging
|
||||
parser.add_argument("--validation-prompt-dir",
|
||||
parser.add_argument("--validation-dataset-file",
|
||||
type=str,
|
||||
help="Directory containing validation prompts")
|
||||
help="Path to unprocessed validation dataset")
|
||||
parser.add_argument("--validation-preprocessed-path",
|
||||
type=str,
|
||||
help="Path to processed validation dataset")
|
||||
parser.add_argument("--validation-sampling-steps",
|
||||
type=str,
|
||||
help="Validation sampling steps")
|
||||
@@ -689,6 +558,9 @@ class TrainingArgs(FastVideoArgs):
|
||||
parser.add_argument("--tracker-project-name",
|
||||
type=str,
|
||||
help="Project name for tracking")
|
||||
parser.add_argument("--wandb-run-name",
|
||||
type=str,
|
||||
help="Run name for wandb")
|
||||
parser.add_argument("--seed",
|
||||
type=int,
|
||||
default=42,
|
||||
@@ -827,4 +699,16 @@ class TrainingArgs(FastVideoArgs):
|
||||
type=str,
|
||||
help="Master weight type")
|
||||
|
||||
# VSA parameters for training with dense to sparse adaption
|
||||
parser.add_argument(
|
||||
"--VSA-decay-rate", # decay rate, how much sparsity you want to decay each step
|
||||
type=float,
|
||||
default=TrainingArgs.VSA_decay_rate,
|
||||
help="VSA decay rate")
|
||||
parser.add_argument(
|
||||
"--VSA-decay-interval-steps", # how many steps for training with current sparsity
|
||||
type=int,
|
||||
default=TrainingArgs.VSA_decay_interval_steps,
|
||||
help="VSA decay interval steps")
|
||||
|
||||
return parser
|
||||
|
||||
@@ -5,15 +5,16 @@ import time
|
||||
from collections import defaultdict
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from typing import Optional
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import torch
|
||||
|
||||
# if TYPE_CHECKING:
|
||||
from fastvideo.v1.attention import AttentionMetadata
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from fastvideo.v1.attention import AttentionMetadata
|
||||
from fastvideo.v1.pipelines import ForwardBatch
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -36,13 +37,13 @@ class ForwardContext:
|
||||
# attn_layers: Dict[str, Any]
|
||||
# TODO: extend to support per-layer dynamic forward context
|
||||
attn_metadata: "AttentionMetadata" # set dynamically for each forward pass
|
||||
forward_batch: Optional[ForwardBatch] = None
|
||||
forward_batch: Optional["ForwardBatch"] = None
|
||||
|
||||
|
||||
_forward_context: Optional[ForwardContext] = None
|
||||
_forward_context: Optional["ForwardContext"] = None
|
||||
|
||||
|
||||
def get_forward_context() -> ForwardContext:
|
||||
def get_forward_context() -> "ForwardContext":
|
||||
"""Get the current forward context."""
|
||||
assert _forward_context is not None, (
|
||||
"Forward context is not set. "
|
||||
@@ -54,7 +55,7 @@ def get_forward_context() -> ForwardContext:
|
||||
@contextmanager
|
||||
def set_forward_context(current_timestep,
|
||||
attn_metadata,
|
||||
forward_batch: Optional[ForwardBatch] = None,
|
||||
forward_batch: Optional["ForwardBatch"] = None,
|
||||
fastvideo_args: Optional[FastVideoArgs] = None):
|
||||
"""A context manager that stores the current forward context,
|
||||
can be attention metadata, etc.
|
||||
|
||||
@@ -5,6 +5,7 @@ from typing import Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from fastvideo.v1.layers.custom_op import CustomOp
|
||||
|
||||
@@ -95,6 +96,22 @@ class ScaleResidual(nn.Module):
|
||||
return residual + x * gate
|
||||
|
||||
|
||||
# adapted from Diffusers: https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/normalization.py
|
||||
# NOTE(will): Needed to match behavior of diffusers and wan2.1 even while using
|
||||
# FSDP's MixedPrecisionPolicy
|
||||
class FP32LayerNorm(nn.LayerNorm):
|
||||
|
||||
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
|
||||
origin_dtype = inputs.dtype
|
||||
return F.layer_norm(
|
||||
inputs.float(),
|
||||
self.normalized_shape,
|
||||
self.weight.float() if self.weight is not None else None,
|
||||
self.bias.float() if self.bias is not None else None,
|
||||
self.eps,
|
||||
).to(origin_dtype)
|
||||
|
||||
|
||||
class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
"""
|
||||
Fused operation that combines:
|
||||
@@ -112,6 +129,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
eps: float = 1e-6,
|
||||
elementwise_affine: bool = False,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
compute_dtype: torch.dtype | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
@@ -121,10 +139,15 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
elif norm_type == "layer":
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
if compute_dtype == torch.float32:
|
||||
self.norm = FP32LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps)
|
||||
else:
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
else:
|
||||
raise NotImplementedError(f"Norm type {norm_type} not implemented")
|
||||
|
||||
@@ -163,18 +186,25 @@ class LayerNormScaleShift(nn.Module):
|
||||
eps: float = 1e-6,
|
||||
elementwise_affine: bool = False,
|
||||
dtype: torch.dtype = torch.float32,
|
||||
compute_dtype: torch.dtype | None = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
self.compute_dtype = compute_dtype
|
||||
if norm_type == "rms":
|
||||
self.norm = RMSNorm(hidden_size,
|
||||
has_weight=elementwise_affine,
|
||||
eps=eps)
|
||||
elif norm_type == "layer":
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
if self.compute_dtype == torch.float32:
|
||||
self.norm = FP32LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps)
|
||||
else:
|
||||
self.norm = nn.LayerNorm(hidden_size,
|
||||
elementwise_affine=elementwise_affine,
|
||||
eps=eps,
|
||||
dtype=dtype)
|
||||
else:
|
||||
raise NotImplementedError(f"Norm type {norm_type} not implemented")
|
||||
|
||||
@@ -182,4 +212,7 @@ class LayerNormScaleShift(nn.Module):
|
||||
scale: torch.Tensor) -> torch.Tensor:
|
||||
"""Apply ln followed by scale and shift in a single fused operation."""
|
||||
normalized = self.norm(x)
|
||||
return normalized * (1.0 + scale) + shift
|
||||
if self.compute_dtype == torch.float32:
|
||||
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
|
||||
else:
|
||||
return normalized * (1.0 + scale) + shift
|
||||
|
||||
@@ -114,7 +114,7 @@ def _info(logger: Logger,
|
||||
|
||||
if (main_process_only and is_main_process) or (local_main_process_only
|
||||
and is_local_main_process):
|
||||
logger.log(logging.INFO, msg, *args, **kwargs)
|
||||
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
|
||||
|
||||
global _warned_local_main_process, _warned_main_process
|
||||
|
||||
@@ -134,7 +134,7 @@ def _info(logger: Logger,
|
||||
_warned_main_process = True
|
||||
|
||||
if not main_process_only and not local_main_process_only:
|
||||
logger.log(logging.INFO, msg, *args, **kwargs)
|
||||
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
|
||||
|
||||
|
||||
class _FastvideoLogger(Logger):
|
||||
|
||||
@@ -6,7 +6,7 @@ import torch
|
||||
from torch import nn
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
# TODO
|
||||
@@ -19,7 +19,7 @@ class BaseDiT(nn.Module, ABC):
|
||||
num_channels_latents: int
|
||||
# always supports torch_sdpa
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = DiTConfig()._supported_attention_backends
|
||||
AttentionBackendEnum, ...] = DiTConfig()._supported_attention_backends
|
||||
|
||||
def __init_subclass__(cls) -> None:
|
||||
required_class_attrs = [
|
||||
@@ -65,7 +65,7 @@ class BaseDiT(nn.Module, ABC):
|
||||
)
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
|
||||
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
|
||||
@@ -85,7 +85,7 @@ class CachableDiT(BaseDiT):
|
||||
num_channels_latents: int
|
||||
# always supports torch_sdpa
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = DiTConfig()._supported_attention_backends
|
||||
AttentionBackendEnum, ...] = DiTConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, config: DiTConfig, **kwargs) -> None:
|
||||
super().__init__(config, **kwargs)
|
||||
|
||||
@@ -23,7 +23,7 @@ from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
|
||||
unpatchify)
|
||||
from fastvideo.v1.models.dits.base import CachableDiT
|
||||
from fastvideo.v1.models.utils import modulate
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class HunyuanRMSNorm(nn.Module):
|
||||
@@ -96,7 +96,8 @@ class MMDoubleStreamBlock(nn.Module):
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
@@ -303,7 +304,8 @@ class MMSingleStreamBlock(nn.Module):
|
||||
num_attention_heads: int,
|
||||
mlp_ratio: float = 4.0,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
prefix: str = "",
|
||||
):
|
||||
super().__init__()
|
||||
@@ -876,8 +878,8 @@ class IndividualTokenRefinerBlock(nn.Module):
|
||||
num_heads=num_attention_heads,
|
||||
head_size=hidden_size // num_attention_heads,
|
||||
# TODO: remove hardcode; remove STA
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA),
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA),
|
||||
)
|
||||
|
||||
def forward(self, x, c):
|
||||
|
||||
@@ -26,7 +26,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
|
||||
get_rotary_pos_embed)
|
||||
from fastvideo.v1.layers.visual_embedding import TimestepEmbedder
|
||||
from fastvideo.v1.models.dits.base import BaseDiT
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class PatchEmbed2D(nn.Module):
|
||||
@@ -139,16 +139,17 @@ class StepVideoRMSNorm(nn.Module):
|
||||
|
||||
class SelfAttention(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
hidden_dim,
|
||||
head_dim,
|
||||
rope_split: Tuple[int, int, int] = (64, 32, 32),
|
||||
bias: bool = False,
|
||||
with_rope: bool = True,
|
||||
with_qk_norm: bool = True,
|
||||
attn_type: str = "torch",
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA)):
|
||||
def __init__(
|
||||
self,
|
||||
hidden_dim,
|
||||
head_dim,
|
||||
rope_split: Tuple[int, int, int] = (64, 32, 32),
|
||||
bias: bool = False,
|
||||
with_rope: bool = True,
|
||||
with_qk_norm: bool = True,
|
||||
attn_type: str = "torch",
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA)):
|
||||
super().__init__()
|
||||
self.head_dim = head_dim
|
||||
self.hidden_dim = hidden_dim
|
||||
@@ -257,7 +258,8 @@ class CrossAttention(nn.Module):
|
||||
head_dim,
|
||||
bias=False,
|
||||
with_qk_norm=True,
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA)
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.head_dim = head_dim
|
||||
|
||||
@@ -14,8 +14,8 @@ from fastvideo.v1.configs.models.dits import WanVideoConfig
|
||||
from fastvideo.v1.configs.sample.wan import WanTeaCacheParams
|
||||
from fastvideo.v1.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.v1.forward_context import get_forward_context
|
||||
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, RMSNorm,
|
||||
ScaleResidual,
|
||||
from fastvideo.v1.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
|
||||
RMSNorm, ScaleResidual,
|
||||
ScaleResidualLayerNormScaleShift)
|
||||
from fastvideo.v1.layers.linear import ReplicatedLinear
|
||||
# from torch.nn import RMSNorm
|
||||
@@ -25,8 +25,11 @@ 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 _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanImageEmbedding(torch.nn.Module):
|
||||
@@ -34,9 +37,9 @@ class WanImageEmbedding(torch.nn.Module):
|
||||
def __init__(self, in_features: int, out_features: int):
|
||||
super().__init__()
|
||||
|
||||
self.norm1 = nn.LayerNorm(in_features)
|
||||
self.norm1 = FP32LayerNorm(in_features)
|
||||
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
|
||||
self.norm2 = nn.LayerNorm(out_features)
|
||||
self.norm2 = FP32LayerNorm(out_features)
|
||||
|
||||
def forward(self,
|
||||
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
@@ -125,8 +128,8 @@ class WanSelfAttention(nn.Module):
|
||||
dropout_rate=0,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
supported_attention_backends=(_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA))
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA))
|
||||
|
||||
def forward(self, x: torch.Tensor, context: torch.Tensor,
|
||||
context_lens: int):
|
||||
@@ -174,7 +177,8 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
window_size=(-1, -1),
|
||||
qk_norm=True,
|
||||
eps=1e-6,
|
||||
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None
|
||||
) -> None:
|
||||
super().__init__(dim, num_heads, window_size, qk_norm, eps,
|
||||
supported_attention_backends)
|
||||
@@ -216,21 +220,22 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
|
||||
class WanTransformerBlock(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
prefix: str = ""):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -261,7 +266,8 @@ class WanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
@@ -281,7 +287,8 @@ class WanTransformerBlock(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -358,21 +365,22 @@ class WanTransformerBlock(nn.Module):
|
||||
|
||||
class WanTransformerBlock_VSA(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
supported_attention_backends: Optional[Tuple[_Backend,
|
||||
...]] = None,
|
||||
prefix: str = ""):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
qk_norm: str = "rms_norm_across_heads",
|
||||
cross_attn_norm: bool = False,
|
||||
eps: float = 1e-6,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
|
||||
...]] = None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.to_q = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_k = ReplicatedLinear(dim, dim, bias=True)
|
||||
self.to_v = ReplicatedLinear(dim, dim, bias=True)
|
||||
@@ -404,7 +412,8 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=True,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
@@ -424,7 +433,8 @@ class WanTransformerBlock_VSA(nn.Module):
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
|
||||
@@ -561,7 +571,8 @@ class WanTransformer3DModel(CachableDiT):
|
||||
norm_type="layer",
|
||||
eps=config.eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32)
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
self.proj_out = nn.Linear(
|
||||
inner_dim, config.out_channels * math.prod(config.patch_size))
|
||||
self.scale_shift_table = nn.Parameter(
|
||||
@@ -569,6 +580,17 @@ class WanTransformer3DModel(CachableDiT):
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
# For type checking
|
||||
self.previous_e0_even = None
|
||||
self.previous_e0_odd = None
|
||||
self.previous_residual_even = None
|
||||
self.previous_residual_odd = None
|
||||
self.is_even = True
|
||||
self.should_calc_even = True
|
||||
self.should_calc_odd = True
|
||||
self.accumulated_rel_l1_distance_even = 0
|
||||
self.accumulated_rel_l1_distance_odd = 0
|
||||
self.cnt = 0
|
||||
self.__post_init__()
|
||||
|
||||
def forward(self,
|
||||
@@ -655,7 +677,7 @@ class WanTransformer3DModel(CachableDiT):
|
||||
# 5. Output norm, projection & unpatchify
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
|
||||
dim=1)
|
||||
hidden_states = self.norm_out(hidden_states.float(), shift, scale)
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
|
||||
|
||||
@@ -8,12 +8,13 @@ from torch import nn
|
||||
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
|
||||
ImageEncoderConfig,
|
||||
TextEncoderConfig)
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
|
||||
|
||||
class TextEncoder(nn.Module, ABC):
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = TextEncoderConfig()._supported_attention_backends
|
||||
AttentionBackendEnum,
|
||||
...] = TextEncoderConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, config: TextEncoderConfig) -> None:
|
||||
super().__init__()
|
||||
@@ -34,13 +35,14 @@ class TextEncoder(nn.Module, ABC):
|
||||
pass
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
|
||||
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
|
||||
class ImageEncoder(nn.Module, ABC):
|
||||
_supported_attention_backends: Tuple[
|
||||
_Backend, ...] = ImageEncoderConfig()._supported_attention_backends
|
||||
AttentionBackendEnum,
|
||||
...] = ImageEncoderConfig()._supported_attention_backends
|
||||
|
||||
def __init__(self, config: ImageEncoderConfig) -> None:
|
||||
super().__init__()
|
||||
@@ -56,5 +58,5 @@ class ImageEncoder(nn.Module, ABC):
|
||||
pass
|
||||
|
||||
@property
|
||||
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
|
||||
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
|
||||
return self._supported_attention_backends
|
||||
|
||||
@@ -81,10 +81,7 @@ def get_hf_config(
|
||||
return config
|
||||
|
||||
|
||||
def get_diffusers_config(
|
||||
model: str,
|
||||
fastvideo_args: Optional[dict] = None,
|
||||
) -> Dict[str, Any]:
|
||||
def get_diffusers_config(model: str, ) -> Dict[str, Any]:
|
||||
"""Gets a configuration for the given diffusers model.
|
||||
|
||||
Args:
|
||||
@@ -105,7 +102,8 @@ def get_diffusers_config(
|
||||
# Load the config directly from the file
|
||||
with open(config_file) as f:
|
||||
config_dict: Dict[str, Any] = json.load(f)
|
||||
|
||||
if "_diffusers_version" in config_dict:
|
||||
config_dict.pop("_diffusers_version")
|
||||
# TODO(will): apply any overrides from inference args
|
||||
return config_dict
|
||||
except Exception as e:
|
||||
|
||||
@@ -15,8 +15,9 @@ from safetensors.torch import load_file as safetensors_load_file
|
||||
from transformers import AutoImageProcessor, AutoTokenizer
|
||||
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
|
||||
|
||||
from fastvideo.v1.configs.models import EncoderConfig
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.hf_transformer_utils import get_diffusers_config
|
||||
from fastvideo.v1.models.loader.fsdp_load import maybe_load_fsdp_model
|
||||
@@ -45,7 +46,7 @@ class ComponentLoader(ABC):
|
||||
Args:
|
||||
model_path: Path to the component model
|
||||
architecture: Architecture of the component model
|
||||
fastvideo_args: Inference arguments
|
||||
fastvideo_args: FastVideoArgs
|
||||
|
||||
Returns:
|
||||
The loaded component
|
||||
@@ -183,9 +184,10 @@ class TextEncoderLoader(ComponentLoader):
|
||||
self,
|
||||
model_config: Any,
|
||||
model: nn.Module,
|
||||
model_path: str,
|
||||
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||
primary_weights = TextEncoderLoader.Source(
|
||||
model_config.model,
|
||||
model_path,
|
||||
prefix="",
|
||||
fall_back_to_pt=getattr(model, "fall_back_to_pt_during_load", True),
|
||||
allow_patterns_overrides=getattr(model, "allow_patterns_overrides",
|
||||
@@ -209,8 +211,7 @@ class TextEncoderLoader(ComponentLoader):
|
||||
# revision=fastvideo_args.revision,
|
||||
# model_override_args=None,
|
||||
# )
|
||||
with open(os.path.join(model_path, "config.json")) as f:
|
||||
model_config = json.load(f)
|
||||
model_config = get_diffusers_config(model=model_path)
|
||||
model_config.pop("_name_or_path", None)
|
||||
model_config.pop("transformers_version", None)
|
||||
model_config.pop("model_type", None)
|
||||
@@ -220,13 +221,17 @@ class TextEncoderLoader(ComponentLoader):
|
||||
|
||||
# @TODO(Wei): Better way to handle this?
|
||||
try:
|
||||
encoder_config = fastvideo_args.text_encoder_configs[0]
|
||||
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[
|
||||
0]
|
||||
encoder_config.update_model_arch(model_config)
|
||||
encoder_precision = fastvideo_args.text_encoder_precisions[0]
|
||||
encoder_precision = fastvideo_args.pipeline_config.text_encoder_precisions[
|
||||
0]
|
||||
except Exception:
|
||||
encoder_config = fastvideo_args.text_encoder_configs[1]
|
||||
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[
|
||||
1]
|
||||
encoder_config.update_model_arch(model_config)
|
||||
encoder_precision = fastvideo_args.text_encoder_precisions[1]
|
||||
encoder_precision = fastvideo_args.pipeline_config.text_encoder_precisions[
|
||||
1]
|
||||
|
||||
target_device = get_torch_device()
|
||||
# TODO(will): add support for other dtypes
|
||||
@@ -235,7 +240,7 @@ class TextEncoderLoader(ComponentLoader):
|
||||
|
||||
def load_model(self,
|
||||
model_path: str,
|
||||
model_config,
|
||||
model_config: EncoderConfig,
|
||||
target_device: torch.device,
|
||||
dtype: str = "fp16"):
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
|
||||
@@ -245,9 +250,8 @@ class TextEncoderLoader(ComponentLoader):
|
||||
model = model_cls(model_config)
|
||||
|
||||
weights_to_load = {name for name, _ in model.named_parameters()}
|
||||
model_config.model = model_path
|
||||
loaded_weights = model.load_weights(
|
||||
self._get_all_weights(model_config, model))
|
||||
self._get_all_weights(model_config, model, model_path))
|
||||
self.counter_after_loading_weights = time.perf_counter()
|
||||
logger.info(
|
||||
"Loading weights took %.2f seconds",
|
||||
@@ -261,7 +265,6 @@ class TextEncoderLoader(ComponentLoader):
|
||||
raise ValueError("Following weights were not initialized from "
|
||||
f"checkpoint: {weights_not_loaded}")
|
||||
|
||||
# TODO(will): add support for training/finetune
|
||||
return model.eval()
|
||||
|
||||
|
||||
@@ -284,13 +287,14 @@ class ImageEncoderLoader(TextEncoderLoader):
|
||||
model_config.pop("model_type", None)
|
||||
logger.info("HF Model config: %s", model_config)
|
||||
|
||||
encoder_config = fastvideo_args.image_encoder_config
|
||||
encoder_config = fastvideo_args.pipeline_config.image_encoder_config
|
||||
encoder_config.update_model_arch(model_config)
|
||||
|
||||
target_device = get_torch_device()
|
||||
# TODO(will): add support for other dtypes
|
||||
return self.load_model(model_path, encoder_config, target_device,
|
||||
fastvideo_args.image_encoder_precision)
|
||||
return self.load_model(
|
||||
model_path, encoder_config, target_device,
|
||||
fastvideo_args.pipeline_config.image_encoder_precision)
|
||||
|
||||
|
||||
class ImageProcessorLoader(ComponentLoader):
|
||||
@@ -332,18 +336,17 @@ class VAELoader(ComponentLoader):
|
||||
def load(self, model_path: str, architecture: str,
|
||||
fastvideo_args: FastVideoArgs):
|
||||
"""Load the VAE based on the model path, architecture, and inference args."""
|
||||
# TODO(will): move this to a constants file
|
||||
config = get_diffusers_config(model=model_path)
|
||||
|
||||
class_name = config.pop("_class_name")
|
||||
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
|
||||
config.pop("_diffusers_version")
|
||||
|
||||
vae_config = fastvideo_args.vae_config
|
||||
vae_config = fastvideo_args.pipeline_config.vae_config
|
||||
vae_config.update_model_arch(config)
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
|
||||
vae = vae_cls(vae_config).to(get_torch_device())
|
||||
with set_default_torch_dtype(PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]):
|
||||
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
vae = vae_cls(vae_config).to(get_torch_device())
|
||||
|
||||
# Find all safetensors files
|
||||
safetensors_list = glob.glob(
|
||||
@@ -355,10 +358,8 @@ class VAELoader(ComponentLoader):
|
||||
loaded = safetensors_load_file(safetensors_list[0])
|
||||
vae.load_state_dict(
|
||||
loaded, strict=False) # We might only load encoder or decoder
|
||||
dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
|
||||
vae = vae.eval().to(dtype)
|
||||
|
||||
return vae
|
||||
return vae.eval()
|
||||
|
||||
|
||||
class TransformerLoader(ComponentLoader):
|
||||
@@ -374,10 +375,9 @@ class TransformerLoader(ComponentLoader):
|
||||
raise ValueError(
|
||||
"Model config does not contain a _class_name attribute. "
|
||||
"Only diffusers format is supported.")
|
||||
config.pop("_diffusers_version")
|
||||
|
||||
# Config from Diffusers supersedes fastvideo's model config
|
||||
dit_config = fastvideo_args.dit_config
|
||||
dit_config = fastvideo_args.pipeline_config.dit_config
|
||||
dit_config.update_model_arch(config)
|
||||
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
|
||||
@@ -391,14 +391,8 @@ class TransformerLoader(ComponentLoader):
|
||||
logger.info("Loading model from %s safetensors files in %s",
|
||||
len(safetensors_list), model_path)
|
||||
|
||||
# initialize_sequence_parallel_group(fastvideo_args.sp_size)
|
||||
if fastvideo_args.training_mode:
|
||||
assert isinstance(
|
||||
fastvideo_args, TrainingArgs
|
||||
), "fastvideo_args must be a TrainingArgs object when training_mode is True"
|
||||
default_dtype = PRECISION_TO_TYPE[fastvideo_args.master_weight_type]
|
||||
else:
|
||||
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
|
||||
default_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.dit_precision]
|
||||
|
||||
# Load the model using FSDP loader
|
||||
logger.info("Loading model from %s, default_dtype: %s", cls_name,
|
||||
@@ -462,15 +456,15 @@ class SchedulerLoader(ComponentLoader):
|
||||
|
||||
class_name = config.pop("_class_name")
|
||||
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
|
||||
config.pop("_diffusers_version")
|
||||
|
||||
scheduler_cls, _ = ModelRegistry.resolve_model_cls(class_name)
|
||||
|
||||
scheduler = scheduler_cls(**config)
|
||||
if fastvideo_args.flow_shift is not None:
|
||||
scheduler.set_shift(fastvideo_args.flow_shift)
|
||||
if fastvideo_args.timesteps_scale is not None:
|
||||
scheduler.set_timesteps_scale(fastvideo_args.timesteps_scale)
|
||||
if fastvideo_args.pipeline_config.flow_shift is not None:
|
||||
scheduler.set_shift(fastvideo_args.pipeline_config.flow_shift)
|
||||
if fastvideo_args.pipeline_config.timesteps_scale is not None:
|
||||
scheduler.set_timesteps_scale(
|
||||
fastvideo_args.pipeline_config.timesteps_scale)
|
||||
return scheduler
|
||||
|
||||
|
||||
@@ -529,7 +523,7 @@ class PipelineComponentLoader:
|
||||
component_model_path: Path to the component model
|
||||
transformers_or_diffusers: Whether the module is from transformers or diffusers
|
||||
architecture: Architecture of the component model
|
||||
fastvideo_args: Inference arguments
|
||||
pipeline_args: Inference arguments
|
||||
|
||||
Returns:
|
||||
The loaded module
|
||||
|
||||
@@ -51,11 +51,6 @@ def auto_attributes(init_func):
|
||||
return wrapper
|
||||
|
||||
|
||||
def set_random_seed(seed: int) -> None:
|
||||
from fastvideo.v1.platforms import current_platform
|
||||
current_platform.seed_everything(seed)
|
||||
|
||||
|
||||
def set_weight_attrs(
|
||||
weight: torch.Tensor,
|
||||
weight_attrs: Optional[Dict[str, Any]],
|
||||
|
||||
@@ -1,8 +1,11 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import os
|
||||
import tempfile
|
||||
from typing import Callable, List, Optional, Tuple, Union
|
||||
from urllib.parse import unquote, urlparse
|
||||
|
||||
import imageio
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
import PIL.ImageOps
|
||||
@@ -86,6 +89,7 @@ def normalize(
|
||||
return 2.0 * images - 1.0
|
||||
|
||||
|
||||
# adapted from diffusers.utils import load_image
|
||||
def load_image(
|
||||
image: Union[str, PIL.Image.Image],
|
||||
convert_method: Optional[Callable[[PIL.Image.Image],
|
||||
@@ -131,6 +135,85 @@ def load_image(
|
||||
return image
|
||||
|
||||
|
||||
# adapted from diffusers.utils import load_video
|
||||
def load_video(
|
||||
video: str,
|
||||
convert_method: Optional[Callable[[List[PIL.Image.Image]],
|
||||
List[PIL.Image.Image]]] = None,
|
||||
) -> List[PIL.Image.Image]:
|
||||
"""
|
||||
Loads `video` to a list of PIL Image.
|
||||
Args:
|
||||
video (`str`):
|
||||
A URL or Path to a video to convert to a list of PIL Image format.
|
||||
convert_method (Callable[[List[PIL.Image.Image]], List[PIL.Image.Image]], *optional*):
|
||||
A conversion method to apply to the video after loading it. When set to `None` the images will be converted
|
||||
to "RGB".
|
||||
Returns:
|
||||
`List[PIL.Image.Image]`:
|
||||
The video as a list of PIL images.
|
||||
"""
|
||||
is_url = video.startswith("http://") or video.startswith("https://")
|
||||
is_file = os.path.isfile(video)
|
||||
was_tempfile_created = False
|
||||
|
||||
if not (is_url or is_file):
|
||||
raise ValueError(
|
||||
f"Incorrect path or URL. URLs must start with `http://` or `https://`, and {video} is not a valid path."
|
||||
)
|
||||
|
||||
if is_url:
|
||||
response = requests.get(video, stream=True)
|
||||
if response.status_code != 200:
|
||||
raise ValueError(
|
||||
f"Failed to download video. Status code: {response.status_code}"
|
||||
)
|
||||
|
||||
parsed_url = urlparse(video)
|
||||
file_name = os.path.basename(unquote(parsed_url.path))
|
||||
|
||||
suffix = os.path.splitext(file_name)[1] or ".mp4"
|
||||
with tempfile.NamedTemporaryFile(suffix=suffix,
|
||||
delete=False) as temp_file:
|
||||
video_path = temp_file.name
|
||||
video_data = response.iter_content(chunk_size=8192)
|
||||
for chunk in video_data:
|
||||
temp_file.write(chunk)
|
||||
|
||||
video = video_path
|
||||
|
||||
pil_images = []
|
||||
if video.endswith(".gif"):
|
||||
gif = PIL.Image.open(video)
|
||||
try:
|
||||
while True:
|
||||
pil_images.append(gif.copy())
|
||||
gif.seek(gif.tell() + 1)
|
||||
except EOFError:
|
||||
pass
|
||||
|
||||
else:
|
||||
try:
|
||||
imageio.plugins.ffmpeg.get_exe()
|
||||
except AttributeError:
|
||||
raise AttributeError(
|
||||
"`Unable to find an ffmpeg installation on your machine. Please install via `pip install imageio-ffmpeg"
|
||||
) from None
|
||||
|
||||
with imageio.get_reader(video) as reader:
|
||||
# Read all frames
|
||||
for frame in reader:
|
||||
pil_images.append(PIL.Image.fromarray(frame))
|
||||
|
||||
if was_tempfile_created:
|
||||
os.remove(video_path)
|
||||
|
||||
if convert_method is not None:
|
||||
pil_images = convert_method(pil_images)
|
||||
|
||||
return pil_images
|
||||
|
||||
|
||||
def get_default_height_width(
|
||||
image: Union[PIL.Image.Image, np.ndarray, torch.Tensor],
|
||||
vae_scale_factor: int,
|
||||
|
||||
@@ -11,7 +11,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.v1.pipelines.lora_pipeline import LoRAPipeline
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import (ForwardBatch,
|
||||
TrainingBatch)
|
||||
from fastvideo.v1.pipelines.pipeline_registry import PipelineRegistry
|
||||
from fastvideo.v1.utils import (maybe_download_model,
|
||||
verify_model_config_and_directory)
|
||||
@@ -49,7 +50,7 @@ def build_pipeline(fastvideo_args: FastVideoArgs) -> PipelineWithLoRA:
|
||||
pipeline_architecture)
|
||||
|
||||
# instantiate the pipeline
|
||||
pipeline = pipeline_cls(model_path, fastvideo_args, config)
|
||||
pipeline = pipeline_cls(model_path, fastvideo_args)
|
||||
logger.info("Pipeline instantiated")
|
||||
|
||||
# pipeline is now initialized and ready to use
|
||||
@@ -63,4 +64,5 @@ __all__ = [
|
||||
"PipelineRegistry",
|
||||
"ForwardBatch",
|
||||
"LoRAPipeline",
|
||||
"TrainingBatch",
|
||||
]
|
||||
|
||||
@@ -8,13 +8,11 @@ This module defines the base class for pipelines that are composed of multiple s
|
||||
import argparse
|
||||
import os
|
||||
from abc import ABC, abstractmethod
|
||||
from copy import deepcopy
|
||||
from typing import Any, Dict, List, Optional, Union, cast
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.configs.pipelines import (PipelineConfig,
|
||||
get_pipeline_config_cls_for_name)
|
||||
from fastvideo.v1.configs.pipelines import PipelineConfig
|
||||
from fastvideo.v1.distributed import (
|
||||
maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
@@ -22,7 +20,7 @@ from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages import PipelineStage
|
||||
from fastvideo.v1.utils import (maybe_download_model, shallow_asdict,
|
||||
from fastvideo.v1.utils import (maybe_download_model,
|
||||
verify_model_config_and_directory)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -46,24 +44,16 @@ class ComposedPipelineBase(ABC):
|
||||
# TODO(will): args should support both inference args and training args
|
||||
def __init__(self,
|
||||
model_path: str,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
config: Optional[Dict[str, Any]] = None,
|
||||
fastvideo_args: Union[FastVideoArgs, TrainingArgs],
|
||||
required_config_modules: Optional[List[str]] = None,
|
||||
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None):
|
||||
"""
|
||||
Initialize the pipeline. After __init__, the pipeline should be ready to
|
||||
use. The pipeline should be stateless and not hold any batch state.
|
||||
"""
|
||||
self.fastvideo_args = fastvideo_args
|
||||
|
||||
if fastvideo_args.training_mode:
|
||||
assert isinstance(fastvideo_args, TrainingArgs)
|
||||
self.training_args = fastvideo_args
|
||||
assert self.training_args is not None
|
||||
else:
|
||||
self.fastvideo_args = fastvideo_args
|
||||
assert self.fastvideo_args is not None
|
||||
|
||||
self.model_path = model_path
|
||||
self.model_path: str = model_path
|
||||
self._stages: List[PipelineStage] = []
|
||||
self._stage_name_mapping: Dict[str, PipelineStage] = {}
|
||||
|
||||
@@ -74,13 +64,6 @@ class ComposedPipelineBase(ABC):
|
||||
raise NotImplementedError(
|
||||
"Subclass must set _required_config_modules")
|
||||
|
||||
if config is None:
|
||||
# Load configuration
|
||||
logger.info("Loading pipeline configuration...")
|
||||
self.config = self._load_config(model_path)
|
||||
else:
|
||||
self.config = config
|
||||
|
||||
maybe_init_distributed_environment_and_model_parallel(
|
||||
fastvideo_args.tp_size, fastvideo_args.sp_size)
|
||||
|
||||
@@ -89,6 +72,8 @@ class ComposedPipelineBase(ABC):
|
||||
self.modules = self.load_modules(fastvideo_args, loaded_modules)
|
||||
|
||||
if fastvideo_args.training_mode:
|
||||
assert isinstance(fastvideo_args, TrainingArgs)
|
||||
self.training_args = fastvideo_args
|
||||
assert self.training_args is not None
|
||||
if self.training_args.log_validation:
|
||||
self.initialize_validation_pipeline(self.training_args)
|
||||
@@ -127,39 +112,16 @@ class ComposedPipelineBase(ABC):
|
||||
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
|
||||
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
|
||||
"""
|
||||
config = None
|
||||
# 1. If users provide a pipeline config, it will override the default pipeline config
|
||||
if isinstance(pipeline_config, PipelineConfig):
|
||||
config = pipeline_config
|
||||
else:
|
||||
config_cls = get_pipeline_config_cls_for_name(model_path)
|
||||
if config_cls is not None:
|
||||
config = config_cls()
|
||||
if isinstance(pipeline_config, str):
|
||||
config.load_from_json(pipeline_config)
|
||||
|
||||
# 2. If users also provide some kwargs, it will override the pipeline config.
|
||||
# The user kwargs shouldn't contain model config parameters!
|
||||
if config is None:
|
||||
logger.warning("No config found for model %s, using default config",
|
||||
model_path)
|
||||
config_args = kwargs
|
||||
else:
|
||||
config_args = shallow_asdict(config)
|
||||
config_args.update(kwargs)
|
||||
|
||||
if args is None or args.inference_mode:
|
||||
fastvideo_args = FastVideoArgs(model_path=model_path, **config_args)
|
||||
|
||||
fastvideo_args.model_path = model_path
|
||||
for key, value in config_args.items():
|
||||
setattr(fastvideo_args, key, value)
|
||||
kwargs['model_path'] = model_path
|
||||
fastvideo_args = FastVideoArgs.from_kwargs(kwargs)
|
||||
else:
|
||||
assert args is not None, "args must be provided for training mode"
|
||||
fastvideo_args = TrainingArgs.from_cli_args(args)
|
||||
# TODO(will): fix this so that its not so ugly
|
||||
fastvideo_args.model_path = model_path
|
||||
for key, value in config_args.items():
|
||||
for key, value in kwargs.items():
|
||||
setattr(fastvideo_args, key, value)
|
||||
|
||||
fastvideo_args.use_cpu_offload = False
|
||||
@@ -170,7 +132,7 @@ class ComposedPipelineBase(ABC):
|
||||
# use FSDP2's MixedPrecisionPolicy to set the precision for the
|
||||
# fwd, bwd, and other operations' precision.
|
||||
# fastvideo_args.precision = fastvideo_args.master_weight_type
|
||||
assert fastvideo_args.master_weight_type == 'fp32', 'only fp32 is supported for training'
|
||||
assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
|
||||
# assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
|
||||
|
||||
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
|
||||
@@ -250,20 +212,21 @@ class ComposedPipelineBase(ABC):
|
||||
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
|
||||
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
|
||||
"""
|
||||
logger.info("Loading pipeline modules from config: %s", self.config)
|
||||
modules_config = deepcopy(self.config)
|
||||
|
||||
model_index = self._load_config(self.model_path)
|
||||
logger.info("Loading pipeline modules from config: %s", model_index)
|
||||
|
||||
# remove keys that are not pipeline modules
|
||||
modules_config.pop("_class_name")
|
||||
modules_config.pop("_diffusers_version")
|
||||
model_index.pop("_class_name")
|
||||
model_index.pop("_diffusers_version")
|
||||
|
||||
# some sanity checks
|
||||
assert len(
|
||||
modules_config
|
||||
model_index
|
||||
) > 1, "model_index.json must contain at least one pipeline module"
|
||||
|
||||
for module_name in self.required_config_modules:
|
||||
if module_name not in modules_config:
|
||||
if module_name not in model_index:
|
||||
raise ValueError(
|
||||
f"model_index.json must contain a {module_name} module")
|
||||
|
||||
@@ -273,7 +236,7 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
modules = {}
|
||||
for module_name, (transformers_or_diffusers,
|
||||
architecture) in modules_config.items():
|
||||
architecture) in model_index.items():
|
||||
if module_name not in required_modules:
|
||||
logger.info("Skipping module %s", module_name)
|
||||
continue
|
||||
@@ -335,3 +298,7 @@ class ComposedPipelineBase(ABC):
|
||||
|
||||
# Return the output
|
||||
return batch
|
||||
|
||||
def train(self) -> None:
|
||||
raise NotImplementedError(
|
||||
"if training_mode is True, the pipeline must implement this method")
|
||||
|
||||
@@ -36,16 +36,17 @@ class LoRAPipeline(ComposedPipelineBase):
|
||||
"transformer"].config.arch_config.exclude_lora_layers
|
||||
|
||||
self.convert_to_lora_layers()
|
||||
if self.fastvideo_args.lora_path is not None:
|
||||
if self.fastvideo_args.pipeline_config.lora_path is not None:
|
||||
self.set_lora_adapter(
|
||||
self.fastvideo_args.lora_nickname, # type: ignore
|
||||
self.fastvideo_args.lora_path)
|
||||
self.fastvideo_args.pipeline_config.
|
||||
lora_nickname, # type: ignore
|
||||
self.fastvideo_args.pipeline_config.lora_path)
|
||||
|
||||
def is_target_layer(self, module_name: str) -> bool:
|
||||
if self.fastvideo_args.lora_target_names is None:
|
||||
if self.fastvideo_args.pipeline_config.lora_target_names is None:
|
||||
return True
|
||||
return any(target_name in module_name
|
||||
for target_name in self.fastvideo_args.lora_target_names)
|
||||
return any(target_name in module_name for target_name in
|
||||
self.fastvideo_args.pipeline_config.lora_target_names)
|
||||
|
||||
def convert_to_lora_layers(self) -> None:
|
||||
"""
|
||||
|
||||
@@ -11,8 +11,10 @@ import pprint
|
||||
from dataclasses import asdict, dataclass, field
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
import PIL.Image
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.attention import AttentionMetadata
|
||||
from fastvideo.v1.configs.sample.teacache import (TeaCacheParams,
|
||||
WanTeaCacheParams)
|
||||
|
||||
@@ -36,6 +38,8 @@ class ForwardBatch:
|
||||
# Image inputs
|
||||
image_path: Optional[str] = None
|
||||
image_embeds: List[torch.Tensor] = field(default_factory=list)
|
||||
pil_image: Optional[PIL.Image.Image] = None
|
||||
preprocessed_image: Optional[torch.Tensor] = None
|
||||
|
||||
# Text inputs
|
||||
prompt: Optional[Union[str, List[str]]] = None
|
||||
@@ -67,6 +71,7 @@ class ForwardBatch:
|
||||
|
||||
# Latent tensors
|
||||
latents: Optional[torch.Tensor] = None
|
||||
raw_latent_shape: Optional[torch.Tensor] = None
|
||||
noise_pred: Optional[torch.Tensor] = None
|
||||
image_latent: Optional[torch.Tensor] = None
|
||||
|
||||
@@ -135,3 +140,38 @@ class ForwardBatch:
|
||||
|
||||
def __str__(self):
|
||||
return pprint.pformat(asdict(self), indent=2, width=120)
|
||||
|
||||
|
||||
@dataclass
|
||||
class TrainingBatch:
|
||||
current_timestep: int = 0
|
||||
current_vsa_sparsity: float = 0.0
|
||||
|
||||
# Dataloader batch outputs
|
||||
latents: Optional[torch.Tensor] = None
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None
|
||||
encoder_attention_mask: Optional[torch.Tensor] = None
|
||||
# i2v
|
||||
# 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
|
||||
infos: Optional[List[Dict[str, Any]]] = None
|
||||
|
||||
# Transformer inputs
|
||||
noisy_model_input: Optional[torch.Tensor] = None
|
||||
timesteps: Optional[torch.Tensor] = None
|
||||
sigmas: Optional[torch.Tensor] = None
|
||||
noise: Optional[torch.Tensor] = None
|
||||
|
||||
attn_metadata: Optional[AttentionMetadata] = None
|
||||
|
||||
# input kwargs
|
||||
input_kwargs: Optional[Dict[str, Any]] = None
|
||||
|
||||
# Training loss
|
||||
loss: torch.Tensor | None = None
|
||||
|
||||
# Training outputs
|
||||
total_loss: float | None = None
|
||||
grad_norm: float | None = None
|
||||
|
||||
+260
-47
@@ -2,19 +2,22 @@
|
||||
import gc
|
||||
import multiprocessing
|
||||
import os
|
||||
from collections import defaultdict
|
||||
from concurrent.futures import ProcessPoolExecutor
|
||||
from itertools import chain
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import numpy as np
|
||||
import PIL.Image
|
||||
import pyarrow as pa
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
from torch.utils.data import DataLoader
|
||||
from torch.utils.data.distributed import DistributedSampler
|
||||
from tqdm import tqdm
|
||||
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
from fastvideo.v1.dataset import getdataset
|
||||
from fastvideo.v1.dataset import ValidationDataset, getdataset
|
||||
from fastvideo.v1.dataset.preprocessing_datasets import PreprocessBatch
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -46,7 +49,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
# Initialize class variables for data sharing
|
||||
self.video_data: Dict[str, Any] = {} # Store video metadata and paths
|
||||
self.latent_data: Dict[str, Any] = {} # Store latent tensors
|
||||
self.preprocess_validation_text(fastvideo_args, args)
|
||||
self.preprocess_validation(fastvideo_args, args)
|
||||
self.preprocess_video_and_text(fastvideo_args, args)
|
||||
|
||||
def get_extra_features(self, valid_data: Dict[str, Any],
|
||||
@@ -58,39 +61,206 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
"""Get the schema fields for the pipeline type. Override in subclasses."""
|
||||
raise NotImplementedError
|
||||
|
||||
def create_record_for_schema(self,
|
||||
preprocess_batch: PreprocessBatch,
|
||||
schema: pa.Schema,
|
||||
strict: bool = False) -> Dict[str, Any]:
|
||||
"""Create a record for the Parquet dataset using a generic schema-based approach.
|
||||
|
||||
Args:
|
||||
preprocess_batch: The batch containing the data to extract
|
||||
schema: PyArrow schema defining the expected fields
|
||||
strict: If True, raises an exception when required fields are missing or unfilled
|
||||
|
||||
Returns:
|
||||
Dictionary record matching the schema
|
||||
|
||||
Raises:
|
||||
ValueError: If strict=True and required fields are missing or unfilled
|
||||
"""
|
||||
record = {}
|
||||
unfilled_fields = []
|
||||
|
||||
for field in schema.names:
|
||||
field_filled = False
|
||||
|
||||
if field.endswith('_bytes'):
|
||||
# Handle binary tensor data - convert numpy array or tensor to bytes
|
||||
tensor_name = field.replace('_bytes', '')
|
||||
tensor_data = getattr(preprocess_batch, tensor_name, None)
|
||||
if tensor_data is not None:
|
||||
try:
|
||||
if hasattr(tensor_data, 'numpy'): # torch tensor
|
||||
record[field] = tensor_data.cpu().numpy().tobytes()
|
||||
field_filled = True
|
||||
elif hasattr(tensor_data, 'tobytes'): # numpy array
|
||||
record[field] = tensor_data.tobytes()
|
||||
field_filled = True
|
||||
else:
|
||||
raise ValueError(
|
||||
f"Unsupported tensor type for field {field}: {type(tensor_data)}"
|
||||
)
|
||||
except Exception as e:
|
||||
if strict:
|
||||
raise ValueError(
|
||||
f"Failed to convert tensor {tensor_name} to bytes: {e}"
|
||||
)
|
||||
record[field] = b'' # Empty bytes for missing data
|
||||
else:
|
||||
record[field] = b'' # Empty bytes for missing data
|
||||
|
||||
elif field.endswith('_shape'):
|
||||
# Handle tensor shape info
|
||||
tensor_name = field.replace('_shape', '')
|
||||
tensor_data = getattr(preprocess_batch, tensor_name, None)
|
||||
if tensor_data is not None and hasattr(tensor_data, 'shape'):
|
||||
record[field] = list(tensor_data.shape)
|
||||
field_filled = True
|
||||
else:
|
||||
record[field] = []
|
||||
|
||||
elif field.endswith('_dtype'):
|
||||
# Handle tensor dtype info
|
||||
tensor_name = field.replace('_dtype', '')
|
||||
tensor_data = getattr(preprocess_batch, tensor_name, None)
|
||||
if tensor_data is not None and hasattr(tensor_data, 'dtype'):
|
||||
record[field] = str(tensor_data.dtype)
|
||||
field_filled = True
|
||||
else:
|
||||
record[field] = 'unknown'
|
||||
|
||||
elif field in ['width', 'height', 'num_frames']:
|
||||
# Handle integer metadata fields
|
||||
value = getattr(preprocess_batch, field, None)
|
||||
if value is not None:
|
||||
try:
|
||||
record[field] = int(value)
|
||||
field_filled = True
|
||||
except (ValueError, TypeError) as e:
|
||||
if strict:
|
||||
raise ValueError(
|
||||
f"Failed to convert field {field} to int: {e}")
|
||||
record[field] = 0
|
||||
else:
|
||||
record[field] = 0
|
||||
|
||||
elif field in ['duration_sec', 'fps']:
|
||||
# Handle float metadata fields
|
||||
# Map schema field names to batch attribute names
|
||||
attr_name = 'duration' if field == 'duration_sec' else field
|
||||
value = getattr(preprocess_batch, attr_name, None)
|
||||
if value is not None:
|
||||
try:
|
||||
record[field] = float(value)
|
||||
field_filled = True
|
||||
except (ValueError, TypeError) as e:
|
||||
if strict:
|
||||
raise ValueError(
|
||||
f"Failed to convert field {field} to float: {e}"
|
||||
)
|
||||
record[field] = 0.0
|
||||
else:
|
||||
record[field] = 0.0
|
||||
|
||||
else:
|
||||
# Handle string fields (id, file_name, caption, media_type, etc.)
|
||||
# Map common schema field names to batch attribute names
|
||||
attr_name = field
|
||||
if field == 'caption':
|
||||
attr_name = 'text'
|
||||
elif field == 'file_name':
|
||||
attr_name = 'path'
|
||||
elif field == 'id':
|
||||
# Generate ID from path if available
|
||||
path_value = getattr(preprocess_batch, 'path', None)
|
||||
if path_value:
|
||||
import os
|
||||
record[field] = os.path.basename(path_value).split(
|
||||
'.')[0]
|
||||
field_filled = True
|
||||
else:
|
||||
record[field] = ""
|
||||
continue
|
||||
elif field == 'media_type':
|
||||
# Determine media type from path
|
||||
path_value = getattr(preprocess_batch, 'path', None)
|
||||
if path_value:
|
||||
record[field] = 'video' if path_value.endswith(
|
||||
'.mp4') else 'image'
|
||||
field_filled = True
|
||||
else:
|
||||
record[field] = ""
|
||||
continue
|
||||
|
||||
value = getattr(preprocess_batch, attr_name, None)
|
||||
if value is not None:
|
||||
record[field] = str(value)
|
||||
field_filled = True
|
||||
else:
|
||||
record[field] = ""
|
||||
|
||||
# Track unfilled fields
|
||||
if not field_filled:
|
||||
unfilled_fields.append(field)
|
||||
|
||||
# Handle strict mode
|
||||
if strict and unfilled_fields:
|
||||
raise ValueError(
|
||||
f"Required fields were not filled: {unfilled_fields}")
|
||||
|
||||
# Log unfilled fields as warning if not in strict mode
|
||||
if unfilled_fields:
|
||||
logger.warning(
|
||||
f"Some fields were not filled and got default values: {unfilled_fields}"
|
||||
)
|
||||
|
||||
return record
|
||||
|
||||
def create_record(
|
||||
self,
|
||||
video_name: str,
|
||||
vae_latent: np.ndarray,
|
||||
text_embedding: np.ndarray,
|
||||
text_attention_mask: np.ndarray,
|
||||
valid_data: Optional[Dict[str, Any]],
|
||||
# text_attention_mask: np.ndarray,
|
||||
valid_data: Dict[str, Any],
|
||||
idx: int,
|
||||
extra_features: Optional[Dict[str, Any]] = None) -> Dict[str, Any]:
|
||||
"""Create a record for the Parquet dataset."""
|
||||
record = {
|
||||
"id": video_name,
|
||||
"vae_latent_bytes": vae_latent.tobytes(),
|
||||
"vae_latent_shape": list(vae_latent.shape),
|
||||
"vae_latent_dtype": str(vae_latent.dtype),
|
||||
"text_embedding_bytes": text_embedding.tobytes(),
|
||||
"text_embedding_shape": list(text_embedding.shape),
|
||||
"text_embedding_dtype": str(text_embedding.dtype),
|
||||
"text_attention_mask_bytes": text_attention_mask.tobytes(),
|
||||
"text_attention_mask_shape": list(text_attention_mask.shape),
|
||||
"text_attention_mask_dtype": str(text_attention_mask.dtype),
|
||||
"file_name": video_name,
|
||||
"caption": valid_data["text"][idx] if valid_data else "",
|
||||
"media_type": "video",
|
||||
"id":
|
||||
video_name,
|
||||
"vae_latent_bytes":
|
||||
vae_latent.tobytes(),
|
||||
"vae_latent_shape":
|
||||
list(vae_latent.shape),
|
||||
"vae_latent_dtype":
|
||||
str(vae_latent.dtype),
|
||||
"text_embedding_bytes":
|
||||
text_embedding.tobytes(),
|
||||
"text_embedding_shape":
|
||||
list(text_embedding.shape),
|
||||
"text_embedding_dtype":
|
||||
str(text_embedding.dtype),
|
||||
"file_name":
|
||||
video_name,
|
||||
"caption":
|
||||
valid_data["text"][idx] if len(valid_data["text"]) > 0 else "",
|
||||
"media_type":
|
||||
"video",
|
||||
"width":
|
||||
valid_data["pixel_values"][idx].shape[-2] if valid_data else 0,
|
||||
valid_data["pixel_values"][idx].shape[-2]
|
||||
if len(valid_data["pixel_values"]) > 0 else 0,
|
||||
"height":
|
||||
valid_data["pixel_values"][idx].shape[-1] if valid_data else 0,
|
||||
valid_data["pixel_values"][idx].shape[-1]
|
||||
if len(valid_data["pixel_values"]) > 0 else 0,
|
||||
"num_frames":
|
||||
vae_latent.shape[1] if len(vae_latent.shape) > 1 else 0,
|
||||
"duration_sec":
|
||||
float(valid_data["duration"][idx]) if valid_data else 0.0,
|
||||
"fps": float(valid_data["fps"][idx]) if valid_data else 0.0,
|
||||
float(valid_data["duration"][idx])
|
||||
if len(valid_data["duration"]) > 0 else 0.0,
|
||||
"fps":
|
||||
float(valid_data["fps"][idx])
|
||||
if len(valid_data["fps"]) > 0 else 0.0,
|
||||
}
|
||||
if extra_features:
|
||||
record.update(extra_features)
|
||||
@@ -103,7 +273,6 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
"combined_parquet_dataset")
|
||||
os.makedirs(combined_parquet_dir, exist_ok=True)
|
||||
local_rank = int(os.getenv("RANK", 0))
|
||||
world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
|
||||
# Get how many samples have already been processed
|
||||
start_idx = 0
|
||||
@@ -114,14 +283,10 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
start_idx += table.num_rows
|
||||
|
||||
# Loading dataset
|
||||
train_dataset = getdataset(args, start_idx=start_idx)
|
||||
sampler = DistributedSampler(train_dataset,
|
||||
rank=local_rank,
|
||||
num_replicas=world_size,
|
||||
shuffle=False)
|
||||
train_dataset = getdataset(args)
|
||||
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
sampler=sampler,
|
||||
batch_size=args.preprocess_video_batch_size,
|
||||
num_workers=args.dataloader_num_workers,
|
||||
)
|
||||
@@ -215,8 +380,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)
|
||||
# text_attention_mask = prompt_attention_mask[idx].cpu().numpy(
|
||||
# ).astype(np.uint8)
|
||||
|
||||
# Get extra features for this sample if needed
|
||||
sample_extra_features = {}
|
||||
@@ -233,7 +398,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
video_name=video_name,
|
||||
vae_latent=vae_latent,
|
||||
text_embedding=text_embedding,
|
||||
text_attention_mask=text_attention_mask,
|
||||
# text_attention_mask=text_attention_mask,
|
||||
valid_data=valid_data,
|
||||
idx=idx,
|
||||
extra_features=sample_extra_features)
|
||||
@@ -285,7 +450,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
num_processed_samples = 0
|
||||
self.all_tables = []
|
||||
|
||||
def preprocess_validation_text(self, fastvideo_args: FastVideoArgs, args):
|
||||
def preprocess_validation(self, fastvideo_args: FastVideoArgs, args):
|
||||
"""Process validation text prompts and save them to parquet files.
|
||||
|
||||
This base implementation handles the common validation text processing logic.
|
||||
@@ -296,22 +461,32 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
"validation_parquet_dataset")
|
||||
os.makedirs(validation_parquet_dir, exist_ok=True)
|
||||
|
||||
with open(args.validation_prompt_txt, encoding="utf-8") as file:
|
||||
lines = file.readlines()
|
||||
prompts = [line.strip() for line in lines]
|
||||
validation_dataset = ValidationDataset(args.validation_dataset_file)
|
||||
|
||||
# Prepare batch data for Parquet dataset
|
||||
batch_data = []
|
||||
sampling_param = SamplingParam.from_pretrained(
|
||||
fastvideo_args.model_path)
|
||||
if sampling_param.negative_prompt:
|
||||
prompts = [sampling_param.negative_prompt] + prompts
|
||||
negative_prompt = {
|
||||
'caption': sampling_param.negative_prompt,
|
||||
'image_path': None,
|
||||
'video_path': None,
|
||||
}
|
||||
validation_iterable = chain([negative_prompt], validation_dataset)
|
||||
else:
|
||||
negative_prompt = None
|
||||
validation_iterable = validation_dataset
|
||||
|
||||
# Add progress bar for validation text preprocessing
|
||||
pbar = tqdm(enumerate(prompts),
|
||||
pbar = tqdm(enumerate(validation_iterable),
|
||||
desc="Processing validation prompts",
|
||||
unit="prompt")
|
||||
for prompt_idx, prompt in pbar:
|
||||
for idx, sample in pbar:
|
||||
with torch.inference_mode():
|
||||
prompt = sample["caption"]
|
||||
is_negative_prompt = idx == 0
|
||||
|
||||
# Text Encoder
|
||||
batch = ForwardBatch(
|
||||
data_type="video",
|
||||
@@ -338,15 +513,44 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
"Shape after removing padding - Embeddings: %s, Mask: %s",
|
||||
text_embedding.shape, text_attention_mask.shape)
|
||||
|
||||
extra_features = {}
|
||||
if not is_negative_prompt:
|
||||
height = sample["height"]
|
||||
width = sample["width"]
|
||||
if "image_path" in sample and "video_path" in sample:
|
||||
raise ValueError(
|
||||
"Only one of image_path or video_path should be provided"
|
||||
)
|
||||
|
||||
if "image" in sample:
|
||||
extra_features = self.preprocess_image(
|
||||
sample["image"], height, width, fastvideo_args)
|
||||
|
||||
if "video" in sample:
|
||||
extra_features = self.preprocess_video(
|
||||
sample["video"], height, width, fastvideo_args)
|
||||
|
||||
# Get extra features for this sample if needed
|
||||
sample_extra_features = {}
|
||||
if extra_features:
|
||||
for key, value in extra_features.items():
|
||||
if isinstance(value, torch.Tensor):
|
||||
sample_extra_features[key] = value.cpu().numpy()
|
||||
else:
|
||||
sample_extra_features[key] = value
|
||||
|
||||
valid_data = defaultdict(list)
|
||||
valid_data["text"] = [prompt]
|
||||
|
||||
# Create record for Parquet dataset
|
||||
record = self.create_record(video_name=file_name,
|
||||
vae_latent=np.array([],
|
||||
dtype=np.float32),
|
||||
text_embedding=text_embedding,
|
||||
text_attention_mask=text_attention_mask,
|
||||
valid_data=None,
|
||||
idx=0,
|
||||
extra_features=None)
|
||||
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)
|
||||
@@ -420,6 +624,15 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
del table
|
||||
gc.collect() # Force garbage collection
|
||||
|
||||
def preprocess_image(self, image: PIL.Image.Image, height: int, width: int,
|
||||
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
return {}
|
||||
|
||||
def preprocess_video(self, video: list[PIL.Image.Image], height: int,
|
||||
width: int,
|
||||
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
return {}
|
||||
|
||||
def _flush_tables(self, num_processed_samples: int, args,
|
||||
combined_parquet_dir: str):
|
||||
"""Flush collected tables to disk."""
|
||||
@@ -8,6 +8,7 @@ using the modular pipeline architecture.
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
import numpy as np
|
||||
import PIL
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
@@ -15,8 +16,13 @@ from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_i2v
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.pipelines.preprocess_pipeline_base import (
|
||||
from fastvideo.v1.models.vision_utils import (get_default_height_width,
|
||||
normalize, numpy_to_pt,
|
||||
pil_to_numpy, resize)
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_base import (
|
||||
BasePreprocessPipeline)
|
||||
from fastvideo.v1.pipelines.stages import ImageEncodingStage, TextEncodingStage
|
||||
|
||||
|
||||
class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
@@ -26,18 +32,75 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
"text_encoder", "tokenizer", "vae", "image_encoder", "image_processor"
|
||||
]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
self.add_stage(stage_name="prompt_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[self.get_module("text_encoder")],
|
||||
tokenizers=[self.get_module("tokenizer")],
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="image_encoding_stage",
|
||||
stage=ImageEncodingStage(
|
||||
image_encoder=self.get_module("image_encoder"),
|
||||
image_processor=self.get_module("image_processor"),
|
||||
))
|
||||
|
||||
def preprocess_image(self, image: PIL.Image.Image, height: int, width: int,
|
||||
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
assert hasattr(
|
||||
self,
|
||||
"image_encoding_stage"), "Image encoding stage must be created"
|
||||
|
||||
batch = ForwardBatch(
|
||||
data_type="video",
|
||||
pil_image=image,
|
||||
)
|
||||
result_batch = self.image_encoding_stage(batch, fastvideo_args)
|
||||
clip_features = result_batch.image_embeds[0]
|
||||
|
||||
# image = self.pil_to_tensor(image)
|
||||
image = self.preprocess(
|
||||
image,
|
||||
vae_scale_factor=self.get_module("vae").spatial_compression_ratio,
|
||||
height=height,
|
||||
width=width)
|
||||
|
||||
return {
|
||||
"clip_feature": clip_features[0],
|
||||
"pil_image": image,
|
||||
}
|
||||
|
||||
def preprocess_video(self, video: list[PIL.Image.Image], height: int,
|
||||
width: int,
|
||||
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
return self.preprocess_image(video[0], height, width, fastvideo_args)
|
||||
|
||||
def get_schema_fields(self) -> List[str]:
|
||||
"""Get the schema fields for I2V pipeline."""
|
||||
return [f.name for f in pyarrow_schema_i2v]
|
||||
|
||||
def get_extra_features(self, valid_data: Dict[str, Any],
|
||||
fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
|
||||
|
||||
# TODO(will): move these to cpu at some point
|
||||
self.get_module("image_encoder").to(get_torch_device())
|
||||
# self.get_module("image_processor").to(get_torch_device())
|
||||
self.get_module("vae").to(get_torch_device())
|
||||
|
||||
features = {}
|
||||
"""Get CLIP features from the first frame of each video."""
|
||||
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
|
||||
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
|
||||
batch_size, _, num_frames, height, width = valid_data[
|
||||
"pixel_values"].shape
|
||||
latent_height = height // self.get_module(
|
||||
"vae").spatial_compression_ratio
|
||||
latent_width = width // self.get_module("vae").spatial_compression_ratio
|
||||
|
||||
processed_images = []
|
||||
# Frame has values between -1 and 1
|
||||
for frame in first_frame:
|
||||
frame = (frame + 1) * 127.5
|
||||
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
|
||||
processed_img = self.get_module("image_processor")(
|
||||
images=frame_pil, return_tensors="pt")
|
||||
@@ -53,25 +116,90 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
clip_features = self.get_module("image_encoder")(**image_inputs)
|
||||
clip_features = clip_features.last_hidden_state
|
||||
|
||||
return {"clip_feature": clip_features}
|
||||
features["clip_feature"] = clip_features
|
||||
"""Get VAE features from the first frame of each video"""
|
||||
video_conditions = []
|
||||
for frame in first_frame:
|
||||
processed_img = frame.to(device="cpu", dtype=torch.float32)
|
||||
processed_img = processed_img.unsqueeze(0).permute(0, 3, 1,
|
||||
2).unsqueeze(2)
|
||||
# (B, H, W, C) -> (B, C, 1, H, W)
|
||||
video_condition = torch.cat([
|
||||
processed_img,
|
||||
processed_img.new_zeros(processed_img.shape[0],
|
||||
processed_img.shape[1], num_frames - 1,
|
||||
height, width)
|
||||
],
|
||||
dim=2)
|
||||
video_condition = video_condition.to(device=get_torch_device(),
|
||||
dtype=torch.float32)
|
||||
video_conditions.append(video_condition)
|
||||
|
||||
video_conditions = torch.cat(video_conditions, dim=0)
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=torch.float32,
|
||||
enabled=True):
|
||||
encoder_outputs = self.get_module("vae").encode(video_conditions)
|
||||
|
||||
latent_condition = encoder_outputs.mean
|
||||
if (hasattr(self.get_module("vae"), "shift_factor")
|
||||
and self.get_module("vae").shift_factor is not None):
|
||||
if isinstance(self.get_module("vae").shift_factor, torch.Tensor):
|
||||
latent_condition -= self.get_module("vae").shift_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition -= self.get_module("vae").shift_factor
|
||||
|
||||
if isinstance(self.get_module("vae").scaling_factor, torch.Tensor):
|
||||
latent_condition = latent_condition * self.get_module(
|
||||
"vae").scaling_factor.to(latent_condition.device,
|
||||
latent_condition.dtype)
|
||||
else:
|
||||
latent_condition = latent_condition * self.get_module(
|
||||
"vae").scaling_factor
|
||||
|
||||
mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
|
||||
latent_width)
|
||||
mask_lat_size[:, :, list(range(1, num_frames))] = 0
|
||||
first_frame_mask = mask_lat_size[:, :, 0:1]
|
||||
first_frame_mask = torch.repeat_interleave(
|
||||
first_frame_mask,
|
||||
dim=2,
|
||||
repeats=self.get_module("vae").temporal_compression_ratio)
|
||||
mask_lat_size = torch.concat(
|
||||
[first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
|
||||
mask_lat_size = mask_lat_size.view(
|
||||
batch_size, -1,
|
||||
self.get_module("vae").temporal_compression_ratio, latent_height,
|
||||
latent_width)
|
||||
mask_lat_size = mask_lat_size.transpose(1, 2)
|
||||
mask_lat_size = mask_lat_size.to(latent_condition.device)
|
||||
|
||||
image_latent = torch.concat([mask_lat_size, latent_condition], dim=1)
|
||||
|
||||
features["first_frame_latent"] = image_latent
|
||||
|
||||
return features
|
||||
|
||||
def create_record(
|
||||
self,
|
||||
video_name: str,
|
||||
vae_latent: np.ndarray,
|
||||
text_embedding: np.ndarray,
|
||||
text_attention_mask: np.ndarray,
|
||||
# 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,
|
||||
text_attention_mask=text_attention_mask,
|
||||
valid_data=valid_data,
|
||||
idx=idx,
|
||||
extra_features=extra_features)
|
||||
record = super().create_record(
|
||||
video_name=video_name,
|
||||
vae_latent=vae_latent,
|
||||
text_embedding=text_embedding,
|
||||
# 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"]
|
||||
@@ -87,7 +215,69 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
|
||||
"clip_feature_dtype": "",
|
||||
})
|
||||
|
||||
if extra_features and "first_frame_latent" in extra_features:
|
||||
first_frame_latent = extra_features["first_frame_latent"]
|
||||
record.update({
|
||||
"first_frame_latent_bytes":
|
||||
first_frame_latent.tobytes(),
|
||||
"first_frame_latent_shape":
|
||||
list(first_frame_latent.shape),
|
||||
"first_frame_latent_dtype":
|
||||
str(first_frame_latent.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"first_frame_latent_bytes": b"",
|
||||
"first_frame_latent_shape": [],
|
||||
"first_frame_latent_dtype": "",
|
||||
})
|
||||
|
||||
if extra_features and "pil_image" in extra_features:
|
||||
pil_image = extra_features["pil_image"]
|
||||
record.update({
|
||||
"pil_image_bytes": pil_image.tobytes(),
|
||||
"pil_image_shape": list(pil_image.shape),
|
||||
"pil_image_dtype": str(pil_image.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"pil_image_bytes": b"",
|
||||
"pil_image_shape": [],
|
||||
"pil_image_dtype": "",
|
||||
})
|
||||
|
||||
return record
|
||||
|
||||
def pil_to_tensor(self, image: PIL.Image.Image) -> torch.Tensor:
|
||||
image = image
|
||||
|
||||
image = np.array(image).astype(np.float32)
|
||||
image = torch.from_numpy(image)
|
||||
return image
|
||||
|
||||
def preprocess(self,
|
||||
image: PIL.Image.Image,
|
||||
vae_scale_factor: int,
|
||||
height: int,
|
||||
width: int,
|
||||
resize_mode: str = "default") -> torch.Tensor:
|
||||
image = [image]
|
||||
|
||||
height, width = get_default_height_width(image[0], vae_scale_factor,
|
||||
height, width)
|
||||
image = [
|
||||
resize(i, height, width, resize_mode=resize_mode) for i in image
|
||||
]
|
||||
image = pil_to_numpy(image) # to np
|
||||
image = numpy_to_pt(image) # to pt
|
||||
|
||||
do_normalize = True
|
||||
if image.min() < 0:
|
||||
do_normalize = False
|
||||
if do_normalize:
|
||||
image = normalize(image)
|
||||
|
||||
return image
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_I2V
|
||||
|
||||
@@ -6,7 +6,7 @@ This module contains an implementation of the T2V Data Preprocessing pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_t2v
|
||||
from fastvideo.v1.pipelines.preprocess_pipeline_base import (
|
||||
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_base import (
|
||||
BasePreprocessPipeline)
|
||||
|
||||
|
||||
|
||||
+39
-38
@@ -1,37 +1,38 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.utils import maybe_download_model, shallow_asdict
|
||||
from fastvideo.v1.distributed import maybe_init_distributed_environment_and_model_parallel, get_world_size
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo import PipelineConfig
|
||||
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_i2v import PreprocessPipeline_I2V
|
||||
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_t2v import PreprocessPipeline_T2V
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.v1.distributed import (
|
||||
get_world_size, maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_i2v import (
|
||||
PreprocessPipeline_I2V)
|
||||
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_t2v import (
|
||||
PreprocessPipeline_T2V)
|
||||
from fastvideo.v1.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
def main(args):
|
||||
args.model_path = maybe_download_model(args.model_path)
|
||||
maybe_init_distributed_environment_and_model_parallel(args.tp_size, args.sp_size)
|
||||
|
||||
def main(args) -> None:
|
||||
args.model_path = maybe_download_model(args.model_path)
|
||||
maybe_init_distributed_environment_and_model_parallel(1, 1)
|
||||
num_gpus = int(os.environ["WORLD_SIZE"])
|
||||
assert num_gpus == 1, "Only support 1 GPU"
|
||||
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
|
||||
kwargs = {
|
||||
"use_cpu_offload": False,
|
||||
"vae_precision": "fp32",
|
||||
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
|
||||
}
|
||||
pipeline_config_args = shallow_asdict(pipeline_config)
|
||||
pipeline_config_args.update(kwargs)
|
||||
fastvideo_args = FastVideoArgs(model_path=args.model_path,
|
||||
num_gpus=get_world_size(),
|
||||
**pipeline_config_args,
|
||||
)
|
||||
pipeline_config.update_config_from_dict(kwargs)
|
||||
fastvideo_args = FastVideoArgs(
|
||||
model_path=args.model_path,
|
||||
num_gpus=get_world_size(),
|
||||
pipeline_config=pipeline_config,
|
||||
)
|
||||
PreprocessPipeline = PreprocessPipeline_I2V if args.preprocess_task == "i2v" else PreprocessPipeline_T2V
|
||||
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
|
||||
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
|
||||
@@ -43,13 +44,14 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--model_type", type=str, default="mochi")
|
||||
parser.add_argument("--data_merge_path", type=str, required=True)
|
||||
parser.add_argument("--validation_prompt_txt", type=str)
|
||||
parser.add_argument("--validation_dataset_file", type=str)
|
||||
parser.add_argument("--num_frames", type=int, default=163)
|
||||
parser.add_argument(
|
||||
"--dataloader_num_workers",
|
||||
type=int,
|
||||
default=1,
|
||||
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
help=
|
||||
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--preprocess_video_batch_size",
|
||||
@@ -63,24 +65,20 @@ if __name__ == "__main__":
|
||||
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,
|
||||
default=256,
|
||||
help="how often to save to parquet files"
|
||||
)
|
||||
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
|
||||
parser.add_argument("--samples_per_file", type=int, default=64)
|
||||
parser.add_argument("--flush_frequency",
|
||||
type=int,
|
||||
default=256,
|
||||
help="how often to save to parquet files")
|
||||
parser.add_argument("--num_latent_t",
|
||||
type=int,
|
||||
default=28,
|
||||
help="Number of latent timesteps.")
|
||||
parser.add_argument("--max_height", type=int, default=480)
|
||||
parser.add_argument("--max_width", type=int, default=848)
|
||||
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--dataset", default="t2v")
|
||||
parser.add_argument("--preprocess_task", type=str, default="t2v")
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
@@ -88,15 +86,18 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--speed_factor", type=float, default=1.0)
|
||||
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
|
||||
# text encoder & vae & diffusion model
|
||||
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
|
||||
parser.add_argument("--text_encoder_name",
|
||||
type=str,
|
||||
default="google/t5-v1_1-xxl")
|
||||
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
|
||||
parser.add_argument("--cfg", type=float, default=0.0)
|
||||
parser.add_argument(
|
||||
"--output_dir",
|
||||
type=str,
|
||||
default=None,
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
help=
|
||||
"The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
main(args)
|
||||
@@ -15,10 +15,16 @@ import torch
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class StageVerificationError(Exception):
|
||||
"""Exception raised when stage verification fails."""
|
||||
pass
|
||||
|
||||
|
||||
class PipelineStage(ABC):
|
||||
"""
|
||||
Abstract base class for all pipeline stages.
|
||||
@@ -28,6 +34,70 @@ class PipelineStage(ABC):
|
||||
for a specific part of the process, such as prompt encoding, latent preparation, etc.
|
||||
"""
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""
|
||||
Verify the input for the stage.
|
||||
|
||||
Example:
|
||||
from fastvideo.v1.pipelines.stages.validators import V, VerificationResult
|
||||
|
||||
def verify_input(self, batch, fastvideo_args):
|
||||
result = VerificationResult()
|
||||
result.add_check("height", batch.height, V.positive_int_divisible(8))
|
||||
result.add_check("width", batch.width, V.positive_int_divisible(8))
|
||||
result.add_check("image_latent", batch.image_latent, V.is_tensor)
|
||||
return result
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
A VerificationResult containing the verification status.
|
||||
|
||||
"""
|
||||
# Default implementation - no verification
|
||||
return VerificationResult()
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""
|
||||
Verify the output for the stage.
|
||||
|
||||
Args:
|
||||
batch: The current batch information.
|
||||
fastvideo_args: The inference arguments.
|
||||
|
||||
Returns:
|
||||
A VerificationResult containing the verification status.
|
||||
"""
|
||||
# Default implementation - no verification
|
||||
return VerificationResult()
|
||||
|
||||
def _run_verification(self, verification_result: VerificationResult,
|
||||
stage_name: str, verification_type: str) -> None:
|
||||
"""
|
||||
Run verification and raise errors if any checks fail.
|
||||
|
||||
Args:
|
||||
verification_result: Results from verify_input or verify_output
|
||||
stage_name: Name of the current stage
|
||||
verification_type: "input" or "output"
|
||||
"""
|
||||
if not verification_result.is_valid():
|
||||
failed_fields = verification_result.get_failed_fields()
|
||||
if failed_fields:
|
||||
# Get detailed failure information
|
||||
detailed_summary = verification_result.get_failure_summary()
|
||||
|
||||
failed_fields_str = ", ".join(failed_fields)
|
||||
error_msg = (
|
||||
f"{verification_type.capitalize()} verification failed for {stage_name}: "
|
||||
f"Failed fields: {failed_fields_str}\n"
|
||||
f"Details: {detailed_summary}")
|
||||
raise StageVerificationError(error_msg)
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
"""Get the device for this stage."""
|
||||
@@ -48,7 +118,7 @@ class PipelineStage(ABC):
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
"""
|
||||
Execute the stage's processing on the batch with optional logging.
|
||||
Execute the stage's processing on the batch with optional verification and logging.
|
||||
Should not be overridden by subclasses.
|
||||
|
||||
Args:
|
||||
@@ -58,34 +128,56 @@ class PipelineStage(ABC):
|
||||
Returns:
|
||||
The updated batch information after this stage's processing.
|
||||
"""
|
||||
# if envs.ENABLE_STAGE_LOGGING:
|
||||
stage_name = self.__class__.__name__
|
||||
|
||||
# Check if verification is enabled (simple approach for prototype)
|
||||
enable_verification = getattr(fastvideo_args,
|
||||
'enable_stage_verification', False)
|
||||
|
||||
if enable_verification:
|
||||
# Pre-execution input verification
|
||||
try:
|
||||
input_result = self.verify_input(batch, fastvideo_args)
|
||||
self._run_verification(input_result, stage_name, "input")
|
||||
except Exception as e:
|
||||
logger.error("Input verification failed for %s: %s", stage_name,
|
||||
str(e))
|
||||
raise
|
||||
|
||||
# Execute the actual stage logic
|
||||
# envs.ENABLE_STAGE_LOGGING
|
||||
if False:
|
||||
self._logger.info("[%s] Starting execution", self._stage_name)
|
||||
self._logger.info("[%s] Starting execution", stage_name)
|
||||
start_time = time.perf_counter()
|
||||
|
||||
try:
|
||||
# Call the actual implementation
|
||||
result = self._call_implementation(batch, fastvideo_args)
|
||||
|
||||
result = self.forward(batch, fastvideo_args)
|
||||
execution_time = time.perf_counter() - start_time
|
||||
self._logger.info("[%s] Execution completed in %s ms",
|
||||
self._stage_name, execution_time * 1000)
|
||||
|
||||
return result
|
||||
stage_name, execution_time * 1000)
|
||||
except Exception as e:
|
||||
execution_time = time.perf_counter() - start_time
|
||||
self._logger.error(
|
||||
"[%s] Error during execution after %s ms: %s",
|
||||
self._stage_name, execution_time * 1000, e)
|
||||
self._logger.error("[%s] Traceback: %s", self._stage_name,
|
||||
"[%s] Error during execution after %s ms: %s", stage_name,
|
||||
execution_time * 1000, e)
|
||||
self._logger.error("[%s] Traceback: %s", stage_name,
|
||||
traceback.format_exc())
|
||||
|
||||
# Re-raise the exception
|
||||
raise
|
||||
else:
|
||||
# Just call the implementation directly if logging is disabled
|
||||
# TODO(will): Also handle backward
|
||||
return self.forward(batch, fastvideo_args)
|
||||
# Direct execution (current behavior)
|
||||
result = self.forward(batch, fastvideo_args)
|
||||
|
||||
if enable_verification:
|
||||
# Post-execution output verification
|
||||
try:
|
||||
output_result = self.verify_output(result, fastvideo_args)
|
||||
self._run_verification(output_result, stage_name, "output")
|
||||
except Exception as e:
|
||||
logger.error("Output verification failed for %s: %s",
|
||||
stage_name, str(e))
|
||||
raise
|
||||
|
||||
return result
|
||||
|
||||
@abstractmethod
|
||||
def forward(
|
||||
|
||||
@@ -9,6 +9,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -69,3 +71,24 @@ class ConditioningStage(PipelineStage):
|
||||
[batch.negative_attention_mask_2, batch.attention_mask_2])
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify conditioning stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("do_classifier_free_guidance",
|
||||
batch.do_classifier_free_guidance, V.bool_value)
|
||||
result.add_check("guidance_scale", batch.guidance_scale,
|
||||
V.positive_float)
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
result.add_check(
|
||||
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
|
||||
not batch.do_classifier_free_guidance or V.list_not_empty(x))
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify conditioning stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
return result
|
||||
|
||||
@@ -11,6 +11,8 @@ from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -27,6 +29,23 @@ class DecodingStage(PipelineStage):
|
||||
def __init__(self, vae) -> None:
|
||||
self.vae: ParallelTiledVAE = vae
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify decoding stage inputs."""
|
||||
result = VerificationResult()
|
||||
# Denoised latents for VAE decoding: [batch_size, channels, frames, height_latents, width_latents]
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify decoding stage outputs."""
|
||||
result = VerificationResult()
|
||||
# Decoded video/images: [batch_size, channels, frames, height, width]
|
||||
result.add_check("output", batch.output, [V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
@@ -54,7 +73,8 @@ class DecodingStage(PipelineStage):
|
||||
image = latents
|
||||
else:
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (vae_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
@@ -77,7 +97,7 @@ class DecodingStage(PipelineStage):
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.vae_tiling:
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
|
||||
@@ -3,15 +3,15 @@
|
||||
Denoising stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
import inspect
|
||||
from typing import Any, Dict, Iterable, List, Optional
|
||||
from typing import Any, Dict, Iterable, Optional
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.v1.attention import get_attn_backend
|
||||
from fastvideo.v1.configs.pipelines.base import STA_Mode
|
||||
from fastvideo.v1.distributed import (get_sp_parallel_rank, get_sp_world_size,
|
||||
get_torch_device, get_world_group)
|
||||
from fastvideo.v1.distributed.communication_op import (
|
||||
@@ -21,19 +21,24 @@ from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.platforms import _Backend
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.v1.platforms import AttentionBackendEnum
|
||||
from fastvideo.v1.utils import dict_to_3d_list
|
||||
|
||||
st_attn_available = False
|
||||
if importlib.util.find_spec("st_attn") is not None:
|
||||
st_attn_available = True
|
||||
try:
|
||||
from fastvideo.v1.attention.backends.sliding_tile_attn import (
|
||||
SlidingTileAttentionBackend)
|
||||
st_attn_available = True
|
||||
except ImportError:
|
||||
st_attn_available = False
|
||||
|
||||
vsa_available = False
|
||||
if importlib.util.find_spec("vsa") is not None:
|
||||
vsa_available = True
|
||||
try:
|
||||
from fastvideo.v1.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionBackend)
|
||||
vsa_available = True
|
||||
except ImportError:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -54,10 +59,11 @@ class DenoisingStage(PipelineStage):
|
||||
self.attn_backend = get_attn_backend(
|
||||
head_size=attn_head_size,
|
||||
dtype=torch.float16, # TODO(will): hack
|
||||
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN,
|
||||
_Backend.VIDEO_SPARSE_ATTN,
|
||||
_Backend.FLASH_ATTN,
|
||||
_Backend.TORCH_SDPA) # hack
|
||||
supported_attention_backends=(
|
||||
AttentionBackendEnum.SLIDING_TILE_ATTN,
|
||||
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA
|
||||
) # hack
|
||||
)
|
||||
|
||||
def forward(
|
||||
@@ -116,20 +122,6 @@ class DenoisingStage(PipelineStage):
|
||||
num_warmup_steps = len(
|
||||
timesteps) - num_inference_steps * self.scheduler.order
|
||||
|
||||
# Create 3D list for mask strategy
|
||||
def dict_to_3d_list(mask_strategy,
|
||||
t_max=50,
|
||||
l_max=60,
|
||||
h_max=24) -> List:
|
||||
result = [[[None for _ in range(h_max)] for _ in range(l_max)]
|
||||
for _ in range(t_max)]
|
||||
if mask_strategy is None:
|
||||
return result
|
||||
for key, value in mask_strategy.items():
|
||||
t, layer, h = map(int, key.split('_'))
|
||||
result[t][layer][h] = value
|
||||
return result
|
||||
|
||||
# Prepare image latents and embeddings for I2V generation
|
||||
image_embeds = batch.image_embeds
|
||||
if len(image_embeds) > 0:
|
||||
@@ -142,7 +134,8 @@ class DenoisingStage(PipelineStage):
|
||||
self.transformer.forward,
|
||||
{
|
||||
"encoder_hidden_states_image": image_embeds,
|
||||
"mask_strategy": dict_to_3d_list(None)
|
||||
"mask_strategy": dict_to_3d_list(
|
||||
None, t_max=50, l_max=60, h_max=24)
|
||||
},
|
||||
)
|
||||
|
||||
@@ -194,13 +187,15 @@ class DenoisingStage(PipelineStage):
|
||||
|
||||
# Prepare inputs for transformer
|
||||
t_expand = t.repeat(latent_model_input.shape[0])
|
||||
guidance_expand = (torch.tensor(
|
||||
[fastvideo_args.embedded_cfg_scale] *
|
||||
latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=get_torch_device(),
|
||||
).to(target_dtype) * 1000.0 if fastvideo_args.embedded_cfg_scale
|
||||
is not None else None)
|
||||
guidance_expand = (
|
||||
torch.tensor(
|
||||
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
|
||||
latent_model_input.shape[0],
|
||||
dtype=torch.float32,
|
||||
device=get_torch_device(),
|
||||
).to(target_dtype) *
|
||||
1000.0 if fastvideo_args.pipeline_config.embedded_cfg_scale
|
||||
is not None else None)
|
||||
|
||||
# Predict noise residual
|
||||
with torch.autocast(device_type="cuda",
|
||||
@@ -301,7 +296,7 @@ class DenoisingStage(PipelineStage):
|
||||
batch.latents = latents
|
||||
|
||||
# Save STA mask search results if needed
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == 'STA_searching':
|
||||
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend and fastvideo_args.STA_mode == STA_Mode.STA_SEARCHING:
|
||||
self.save_sta_search_results(batch)
|
||||
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
@@ -399,7 +394,7 @@ class DenoisingStage(PipelineStage):
|
||||
raise NotImplementedError(
|
||||
"STA mask search/tuning is not supported for this resolution")
|
||||
|
||||
if STA_mode == "STA_searching" or STA_mode == "STA_tuning" or STA_mode == "STA_tuning_cfg":
|
||||
if STA_mode == STA_Mode.STA_SEARCHING or STA_mode == STA_Mode.STA_TUNING or STA_mode == STA_Mode.STA_TUNING_CFG:
|
||||
size = (batch.width, batch.height)
|
||||
if size == (1280, 768):
|
||||
# TODO: make it configurable
|
||||
@@ -421,18 +416,18 @@ class DenoisingStage(PipelineStage):
|
||||
layer_num += self.transformer.config.num_single_layers
|
||||
head_num = self.transformer.config.num_attention_heads
|
||||
|
||||
if STA_mode == "STA_searching":
|
||||
if STA_mode == STA_Mode.STA_SEARCHING:
|
||||
STA_param = configure_sta(
|
||||
mode='STA_searching',
|
||||
mode=STA_Mode.STA_SEARCHING,
|
||||
layer_num=layer_num,
|
||||
head_num=head_num,
|
||||
time_step_num=timesteps_num,
|
||||
mask_candidates=sparse_mask_candidates_searching +
|
||||
full_mask, # last is full mask; Can add more sparse masks while keep last one as full mask
|
||||
)
|
||||
elif STA_mode == 'STA_tuning':
|
||||
elif STA_mode == STA_Mode.STA_TUNING:
|
||||
STA_param = configure_sta(
|
||||
mode='STA_tuning',
|
||||
mode=STA_Mode.STA_TUNING,
|
||||
layer_num=layer_num,
|
||||
head_num=head_num,
|
||||
time_step_num=timesteps_num,
|
||||
@@ -445,9 +440,9 @@ class DenoisingStage(PipelineStage):
|
||||
save_dir=
|
||||
f'output/mask_search_strategy_{size[0]}x{size[1]}/', # Custom save directory
|
||||
timesteps=timesteps_num)
|
||||
elif STA_mode == 'STA_tuning_cfg':
|
||||
elif STA_mode == STA_Mode.STA_TUNING_CFG:
|
||||
STA_param = configure_sta(
|
||||
mode='STA_tuning_cfg',
|
||||
mode=STA_Mode.STA_TUNING_CFG,
|
||||
layer_num=layer_num,
|
||||
head_num=head_num,
|
||||
time_step_num=timesteps_num,
|
||||
@@ -460,12 +455,12 @@ class DenoisingStage(PipelineStage):
|
||||
skip_time_steps=skip_time_steps,
|
||||
save_dir=f'output/mask_search_strategy_{size[0]}x{size[1]}/',
|
||||
timesteps=timesteps_num)
|
||||
elif STA_mode == 'STA_inference':
|
||||
elif STA_mode == STA_Mode.STA_INFERENCE:
|
||||
import fastvideo.v1.envs as envs
|
||||
config_file = envs.FASTVIDEO_ATTENTION_CONFIG
|
||||
if config_file is None:
|
||||
raise ValueError("FASTVIDEO_ATTENTION_CONFIG is not set")
|
||||
STA_param = configure_sta(mode='STA_inference',
|
||||
STA_param = configure_sta(mode=STA_Mode.STA_INFERENCE,
|
||||
layer_num=layer_num,
|
||||
head_num=head_num,
|
||||
time_step_num=timesteps_num,
|
||||
@@ -514,3 +509,37 @@ class DenoisingStage(PipelineStage):
|
||||
mask_strategies=sparse_mask_candidates_searching,
|
||||
output_dir=f'output/mask_search_result_neg_{size[0]}x{size[1]}/'
|
||||
)
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify denoising stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("timesteps", batch.timesteps,
|
||||
[V.is_tensor, V.min_dims(1)])
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
|
||||
result.add_check("image_embeds", batch.image_embeds, V.is_list)
|
||||
result.add_check("image_latent", batch.image_latent,
|
||||
V.none_or_tensor_with_dims(5))
|
||||
result.add_check("num_inference_steps", batch.num_inference_steps,
|
||||
V.positive_int)
|
||||
result.add_check("guidance_scale", batch.guidance_scale,
|
||||
V.positive_float)
|
||||
result.add_check("eta", batch.eta, V.non_negative_float)
|
||||
result.add_check("generator", batch.generator,
|
||||
V.generator_or_list_generators)
|
||||
result.add_check("do_classifier_free_guidance",
|
||||
batch.do_classifier_free_guidance, V.bool_value)
|
||||
result.add_check(
|
||||
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
|
||||
not batch.do_classifier_free_guidance or V.list_not_empty(x))
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify denoising stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
@@ -12,10 +12,12 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
|
||||
from fastvideo.v1.models.vision_utils import (get_default_height_width,
|
||||
load_image, normalize,
|
||||
numpy_to_pt, pil_to_numpy, resize)
|
||||
normalize, numpy_to_pt,
|
||||
pil_to_numpy, resize)
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import V # Import validators
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
from fastvideo.v1.utils import PRECISION_TO_TYPE
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -49,22 +51,26 @@ class EncodingStage(PipelineStage):
|
||||
"""
|
||||
self.vae = self.vae.to(get_torch_device())
|
||||
|
||||
image_path = batch.image_path
|
||||
# TODO(will): remove this once we add input/output validation for stages
|
||||
if image_path is None:
|
||||
raise ValueError("Image Path must be provided")
|
||||
assert batch.height is not None
|
||||
assert batch.width is not None
|
||||
latent_height = batch.height // self.vae.spatial_compression_ratio
|
||||
latent_width = batch.width // self.vae.spatial_compression_ratio
|
||||
|
||||
image = load_image(image_path)
|
||||
image = self.preprocess(
|
||||
image,
|
||||
vae_scale_factor=self.vae.spatial_compression_ratio,
|
||||
height=batch.height,
|
||||
width=batch.width).to(get_torch_device(), dtype=torch.float32)
|
||||
image = image.unsqueeze(2)
|
||||
image = batch.preprocessed_image
|
||||
# TODO(will)
|
||||
if image is None:
|
||||
assert batch.pil_image is not None
|
||||
image = batch.pil_image
|
||||
image = self.preprocess(
|
||||
image,
|
||||
vae_scale_factor=self.vae.spatial_compression_ratio,
|
||||
height=batch.height,
|
||||
width=batch.width).to(get_torch_device(), dtype=torch.float32)
|
||||
|
||||
image = image.unsqueeze(2)
|
||||
else:
|
||||
image = image.transpose(1, 2)
|
||||
logger.info("image: %s", image.shape)
|
||||
video_condition = torch.cat([
|
||||
image,
|
||||
image.new_zeros(image.shape[0], image.shape[1],
|
||||
@@ -75,7 +81,8 @@ class EncodingStage(PipelineStage):
|
||||
dtype=torch.float32)
|
||||
|
||||
# Setup VAE precision
|
||||
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
vae_autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
@@ -83,7 +90,7 @@ class EncodingStage(PipelineStage):
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=vae_autocast_enabled):
|
||||
if fastvideo_args.vae_tiling:
|
||||
if fastvideo_args.pipeline_config.vae_tiling:
|
||||
self.vae.enable_tiling()
|
||||
# if fastvideo_args.vae_sp:
|
||||
# self.vae.enable_parallel()
|
||||
@@ -94,7 +101,7 @@ class EncodingStage(PipelineStage):
|
||||
generator = batch.generator
|
||||
if generator is None:
|
||||
raise ValueError("Generator must be provided")
|
||||
latent_condition = self.retrieve_latents(encoder_output, generator[0])
|
||||
latent_condition = self.retrieve_latents(encoder_output, generator)
|
||||
|
||||
# Apply shifting if needed
|
||||
if (hasattr(self.vae, "shift_factor")
|
||||
@@ -173,3 +180,23 @@ class EncodingStage(PipelineStage):
|
||||
image = normalize(image)
|
||||
|
||||
return image
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify encoding stage inputs."""
|
||||
result = VerificationResult()
|
||||
# result.add_check("pil_image", batch.pil_image)
|
||||
result.add_check("height", batch.height, V.positive_int)
|
||||
result.add_check("width", batch.width, V.positive_int)
|
||||
result.add_check("generator", batch.generator,
|
||||
V.generator_or_list_generators)
|
||||
result.add_check("num_frames", batch.num_frames, V.positive_int)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify encoding stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("image_latent", batch.image_latent,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
return result
|
||||
|
||||
@@ -11,9 +11,10 @@ from fastvideo.v1.distributed import get_torch_device
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vision_utils import load_image
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -56,7 +57,7 @@ class ImageEncodingStage(PipelineStage):
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
self.image_encoder = self.image_encoder.to(get_torch_device())
|
||||
|
||||
image = load_image(batch.image_path)
|
||||
image = batch.pil_image
|
||||
|
||||
image_inputs = self.image_processor(
|
||||
images=image, return_tensors="pt").to(get_torch_device())
|
||||
@@ -71,3 +72,19 @@ class ImageEncodingStage(PipelineStage):
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify image encoding stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("pil_image", batch.pil_image, V.not_none)
|
||||
result.add_check("image_embeds", batch.image_embeds, V.is_list)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify image encoding stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("image_embeds", batch.image_embeds,
|
||||
V.list_of_tensors_dims(3))
|
||||
return result
|
||||
|
||||
@@ -7,11 +7,17 @@ import torch
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.vision_utils import load_image
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import (StageValidators,
|
||||
VerificationResult)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
# Alias for convenience
|
||||
V = StageValidators
|
||||
|
||||
|
||||
class InputValidationStage(PipelineStage):
|
||||
"""
|
||||
@@ -86,4 +92,37 @@ class InputValidationStage(PipelineStage):
|
||||
f"Guidance scale must be positive, but got {batch.guidance_scale}"
|
||||
)
|
||||
|
||||
# for i2v, get image from image_path
|
||||
if batch.image_path is not None:
|
||||
image = load_image(batch.image_path)
|
||||
batch.pil_image = image
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify input validation stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("seed", batch.seed, [V.not_none, V.positive_int])
|
||||
result.add_check("num_videos_per_prompt", batch.num_videos_per_prompt,
|
||||
V.positive_int)
|
||||
result.add_check(
|
||||
"prompt_or_embeds", None, lambda _: V.string_or_list_strings(
|
||||
batch.prompt) or V.list_not_empty(batch.prompt_embeds))
|
||||
result.add_check("height", batch.height, V.positive_int)
|
||||
result.add_check("width", batch.width, V.positive_int)
|
||||
result.add_check("num_inference_steps", batch.num_inference_steps,
|
||||
V.positive_int)
|
||||
result.add_check(
|
||||
"guidance_scale", batch.guidance_scale, lambda x: not batch.
|
||||
do_classifier_free_guidance or V.positive_float(x))
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify input validation stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("seeds", batch.seeds, V.list_not_empty)
|
||||
result.add_check("generator", batch.generator,
|
||||
V.generator_or_list_generators)
|
||||
return result
|
||||
|
||||
@@ -2,6 +2,7 @@
|
||||
"""
|
||||
Latent preparation stage for diffusion pipelines.
|
||||
"""
|
||||
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.v1.distributed import get_torch_device
|
||||
@@ -9,6 +10,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -75,10 +78,10 @@ class LatentPreparationStage(PipelineStage):
|
||||
batch_size,
|
||||
self.transformer.num_channels_latents,
|
||||
num_frames,
|
||||
height //
|
||||
fastvideo_args.vae_config.arch_config.spatial_compression_ratio,
|
||||
width //
|
||||
fastvideo_args.vae_config.arch_config.spatial_compression_ratio,
|
||||
height // fastvideo_args.pipeline_config.vae_config.arch_config.
|
||||
spatial_compression_ratio,
|
||||
width // fastvideo_args.pipeline_config.vae_config.arch_config.
|
||||
spatial_compression_ratio,
|
||||
)
|
||||
|
||||
# Validate generator if it's a list
|
||||
@@ -103,6 +106,7 @@ class LatentPreparationStage(PipelineStage):
|
||||
|
||||
# Update batch with prepared latents
|
||||
batch.latents = latents
|
||||
batch.raw_latent_shape = latents.shape
|
||||
|
||||
return batch
|
||||
|
||||
@@ -119,10 +123,38 @@ class LatentPreparationStage(PipelineStage):
|
||||
The batch with adjusted video length.
|
||||
"""
|
||||
video_length = batch.num_frames
|
||||
use_temporal_scaling_frames = fastvideo_args.vae_config.use_temporal_scaling_frames
|
||||
use_temporal_scaling_frames = fastvideo_args.pipeline_config.vae_config.use_temporal_scaling_frames
|
||||
if use_temporal_scaling_frames:
|
||||
temporal_scale_factor = fastvideo_args.vae_config.arch_config.temporal_compression_ratio
|
||||
temporal_scale_factor = fastvideo_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
latent_num_frames = (video_length - 1) // temporal_scale_factor + 1
|
||||
else: # stepvideo only
|
||||
latent_num_frames = video_length // 17 * 3
|
||||
return latent_num_frames
|
||||
return int(latent_num_frames)
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify latent preparation stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check(
|
||||
"prompt_or_embeds", None, lambda _: V.string_or_list_strings(
|
||||
batch.prompt) or V.list_not_empty(batch.prompt_embeds))
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds,
|
||||
V.list_of_tensors)
|
||||
result.add_check("num_videos_per_prompt", batch.num_videos_per_prompt,
|
||||
V.positive_int)
|
||||
result.add_check("generator", batch.generator,
|
||||
V.generator_or_list_generators)
|
||||
result.add_check("num_frames", batch.num_frames, V.positive_int)
|
||||
result.add_check("height", batch.height, V.positive_int)
|
||||
result.add_check("width", batch.width, V.positive_int)
|
||||
result.add_check("latents", batch.latents, V.none_or_tensor)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify latent preparation stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("latents", batch.latents,
|
||||
[V.is_tensor, V.with_dims(5)])
|
||||
result.add_check("raw_latent_shape", batch.raw_latent_shape, V.is_tuple)
|
||||
return result
|
||||
|
||||
@@ -1,10 +1,14 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -29,9 +33,9 @@ class StepvideoPromptEncodingStage(PipelineStage):
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args) -> ForwardBatch:
|
||||
|
||||
prompts = [batch.prompt + fastvideo_args.pos_magic]
|
||||
prompts = [batch.prompt + fastvideo_args.pipeline_config.pos_magic]
|
||||
bs = len(prompts)
|
||||
prompts += [fastvideo_args.neg_magic] * bs
|
||||
prompts += [fastvideo_args.pipeline_config.neg_magic] * bs
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
y, y_mask = self.stepllm(prompts)
|
||||
clip_emb, _ = self.clip(prompts)
|
||||
@@ -47,3 +51,29 @@ class StepvideoPromptEncodingStage(PipelineStage):
|
||||
batch.clip_embedding_pos = pos_clip
|
||||
batch.clip_embedding_neg = neg_clip
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify stepvideo encoding stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("prompt", batch.prompt, V.string_not_empty)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify stepvideo encoding stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds,
|
||||
[V.is_tensor, V.with_dims(3)])
|
||||
result.add_check("negative_prompt_embeds", batch.negative_prompt_embeds,
|
||||
[V.is_tensor, V.with_dims(3)])
|
||||
result.add_check("prompt_attention_mask", batch.prompt_attention_mask,
|
||||
[V.is_tensor, V.with_dims(2)])
|
||||
result.add_check("negative_attention_mask",
|
||||
batch.negative_attention_mask,
|
||||
[V.is_tensor, V.with_dims(2)])
|
||||
result.add_check("clip_embedding_pos", batch.clip_embedding_pos,
|
||||
[V.is_tensor, V.with_dims(2)])
|
||||
result.add_check("clip_embedding_neg", batch.clip_embedding_neg,
|
||||
[V.is_tensor, V.with_dims(2)])
|
||||
return result
|
||||
|
||||
@@ -12,6 +12,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = (__name__)
|
||||
|
||||
@@ -53,13 +55,13 @@ class TextEncodingStage(PipelineStage):
|
||||
"""
|
||||
assert len(self.tokenizers) == len(self.text_encoders)
|
||||
assert len(self.text_encoders) == len(
|
||||
fastvideo_args.text_encoder_configs)
|
||||
fastvideo_args.pipeline_config.text_encoder_configs)
|
||||
|
||||
for tokenizer, text_encoder, encoder_config, preprocess_func, postprocess_func in zip(
|
||||
self.tokenizers, self.text_encoders,
|
||||
fastvideo_args.text_encoder_configs,
|
||||
fastvideo_args.preprocess_text_funcs,
|
||||
fastvideo_args.postprocess_text_funcs):
|
||||
fastvideo_args.pipeline_config.text_encoder_configs,
|
||||
fastvideo_args.pipeline_config.preprocess_text_funcs,
|
||||
fastvideo_args.pipeline_config.postprocess_text_funcs):
|
||||
if fastvideo_args.use_cpu_offload:
|
||||
text_encoder = text_encoder.to(get_torch_device())
|
||||
|
||||
@@ -113,3 +115,30 @@ class TextEncodingStage(PipelineStage):
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify text encoding stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("prompt", batch.prompt, V.string_or_list_strings)
|
||||
result.add_check(
|
||||
"negative_prompt", batch.negative_prompt, lambda x: not batch.
|
||||
do_classifier_free_guidance or V.string_not_empty(x))
|
||||
result.add_check("do_classifier_free_guidance",
|
||||
batch.do_classifier_free_guidance, V.bool_value)
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds, V.is_list)
|
||||
result.add_check("negative_prompt_embeds", batch.negative_prompt_embeds,
|
||||
V.none_or_list)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify text encoding stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("prompt_embeds", batch.prompt_embeds,
|
||||
V.list_of_tensors_min_dims(2))
|
||||
result.add_check(
|
||||
"negative_prompt_embeds", batch.negative_prompt_embeds,
|
||||
lambda x: not batch.do_classifier_free_guidance or V.
|
||||
list_of_tensors_with_min_dims(x, 2))
|
||||
return result
|
||||
|
||||
@@ -12,6 +12,8 @@ from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.v1.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.v1.pipelines.stages.validators import StageValidators as V
|
||||
from fastvideo.v1.pipelines.stages.validators import VerificationResult
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -95,3 +97,22 @@ class TimestepPreparationStage(PipelineStage):
|
||||
batch.timesteps = timesteps
|
||||
|
||||
return batch
|
||||
|
||||
def verify_input(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify timestep preparation stage inputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("num_inference_steps", batch.num_inference_steps,
|
||||
V.positive_int)
|
||||
result.add_check("timesteps", batch.timesteps, V.none_or_tensor)
|
||||
result.add_check("sigmas", batch.sigmas, V.none_or_list)
|
||||
result.add_check("n_tokens", batch.n_tokens, V.none_or_positive_int)
|
||||
return result
|
||||
|
||||
def verify_output(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> VerificationResult:
|
||||
"""Verify timestep preparation stage outputs."""
|
||||
result = VerificationResult()
|
||||
result.add_check("timesteps", batch.timesteps,
|
||||
[V.is_tensor, V.with_dims(1)])
|
||||
return result
|
||||
|
||||
@@ -0,0 +1,486 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Common validators for pipeline stage verification.
|
||||
|
||||
This module provides reusable validation functions that can be used across
|
||||
all pipeline stages for input/output verification.
|
||||
"""
|
||||
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
class StageValidators:
|
||||
"""Common validators for pipeline stages."""
|
||||
|
||||
@staticmethod
|
||||
def not_none(value: Any) -> bool:
|
||||
"""Check if value is not None."""
|
||||
return value is not None
|
||||
|
||||
@staticmethod
|
||||
def positive_int(value: Any) -> bool:
|
||||
"""Check if value is a positive integer."""
|
||||
return isinstance(value, int) and value > 0
|
||||
|
||||
@staticmethod
|
||||
def positive_float(value: Any) -> bool:
|
||||
"""Check if value is a positive float."""
|
||||
return isinstance(value, (int, float)) and value > 0
|
||||
|
||||
@staticmethod
|
||||
def non_negative_float(value: Any) -> bool:
|
||||
"""Check if value is a non-negative float."""
|
||||
return isinstance(value, (int, float)) and value >= 0
|
||||
|
||||
@staticmethod
|
||||
def divisible_by(value: Any, divisor: int) -> bool:
|
||||
"""Check if value is divisible by divisor."""
|
||||
return value is not None and isinstance(value,
|
||||
int) and value % divisor == 0
|
||||
|
||||
@staticmethod
|
||||
def is_tensor(value: Any) -> bool:
|
||||
"""Check if value is a torch tensor and doesn't contain NaN values."""
|
||||
if not isinstance(value, torch.Tensor):
|
||||
return False
|
||||
return not torch.isnan(value).any().item()
|
||||
|
||||
@staticmethod
|
||||
def tensor_with_dims(value: Any, dims: int) -> bool:
|
||||
"""Check if value is a tensor with specific dimensions and no NaN values."""
|
||||
if not isinstance(value, torch.Tensor):
|
||||
return False
|
||||
if value.dim() != dims:
|
||||
return False
|
||||
return not torch.isnan(value).any().item()
|
||||
|
||||
@staticmethod
|
||||
def tensor_min_dims(value: Any, min_dims: int) -> bool:
|
||||
"""Check if value is a tensor with at least min_dims dimensions and no NaN values."""
|
||||
if not isinstance(value, torch.Tensor):
|
||||
return False
|
||||
if value.dim() < min_dims:
|
||||
return False
|
||||
return not torch.isnan(value).any().item()
|
||||
|
||||
@staticmethod
|
||||
def tensor_shape_matches(value: Any, expected_shape: tuple) -> bool:
|
||||
"""Check if tensor shape matches expected shape (None for any size) and no NaN values."""
|
||||
if not isinstance(value, torch.Tensor):
|
||||
return False
|
||||
if len(value.shape) != len(expected_shape):
|
||||
return False
|
||||
for actual, expected in zip(value.shape, expected_shape):
|
||||
if expected is not None and actual != expected:
|
||||
return False
|
||||
return not torch.isnan(value).any().item()
|
||||
|
||||
@staticmethod
|
||||
def list_not_empty(value: Any) -> bool:
|
||||
"""Check if value is a non-empty list."""
|
||||
return isinstance(value, list) and len(value) > 0
|
||||
|
||||
@staticmethod
|
||||
def list_length(value: Any, length: int) -> bool:
|
||||
"""Check if list has specific length."""
|
||||
return isinstance(value, list) and len(value) == length
|
||||
|
||||
@staticmethod
|
||||
def list_min_length(value: Any, min_length: int) -> bool:
|
||||
"""Check if list has at least min_length items."""
|
||||
return isinstance(value, list) and len(value) >= min_length
|
||||
|
||||
@staticmethod
|
||||
def string_not_empty(value: Any) -> bool:
|
||||
"""Check if value is a non-empty string."""
|
||||
return isinstance(value, str) and len(value.strip()) > 0
|
||||
|
||||
@staticmethod
|
||||
def string_or_list_strings(value: Any) -> bool:
|
||||
"""Check if value is a string or list of strings."""
|
||||
if isinstance(value, str):
|
||||
return True
|
||||
if isinstance(value, list):
|
||||
return all(isinstance(item, str) for item in value)
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def bool_value(value: Any) -> bool:
|
||||
"""Check if value is a boolean."""
|
||||
return isinstance(value, bool)
|
||||
|
||||
@staticmethod
|
||||
def generator_or_list_generators(value: Any) -> bool:
|
||||
"""Check if value is a Generator or list of Generators."""
|
||||
if isinstance(value, torch.Generator):
|
||||
return True
|
||||
if isinstance(value, list):
|
||||
return all(isinstance(item, torch.Generator) for item in value)
|
||||
return False
|
||||
|
||||
@staticmethod
|
||||
def is_list(value: Any) -> bool:
|
||||
"""Check if value is a list (can be empty)."""
|
||||
return isinstance(value, list)
|
||||
|
||||
@staticmethod
|
||||
def is_tuple(value: Any) -> bool:
|
||||
"""Check if value is a tuple."""
|
||||
return isinstance(value, tuple)
|
||||
|
||||
@staticmethod
|
||||
def none_or_tensor(value: Any) -> bool:
|
||||
"""Check if value is None or a tensor without NaN values."""
|
||||
if value is None:
|
||||
return True
|
||||
if not isinstance(value, torch.Tensor):
|
||||
return False
|
||||
return not torch.isnan(value).any().item()
|
||||
|
||||
@staticmethod
|
||||
def list_of_tensors_with_dims(value: Any, dims: int) -> bool:
|
||||
"""Check if value is a non-empty list where all items are tensors with specific dimensions and no NaN values."""
|
||||
if not isinstance(value, list) or len(value) == 0:
|
||||
return False
|
||||
for item in value:
|
||||
if not isinstance(item, torch.Tensor):
|
||||
return False
|
||||
if item.dim() != dims:
|
||||
return False
|
||||
if torch.isnan(item).any().item():
|
||||
return False
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def list_of_tensors(value: Any) -> bool:
|
||||
"""Check if value is a non-empty list where all items are tensors without NaN values."""
|
||||
if not isinstance(value, list) or len(value) == 0:
|
||||
return False
|
||||
for item in value:
|
||||
if not isinstance(item, torch.Tensor):
|
||||
return False
|
||||
if torch.isnan(item).any().item():
|
||||
return False
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def list_of_tensors_with_min_dims(value: Any, min_dims: int) -> bool:
|
||||
"""Check if value is a non-empty list where all items are tensors with at least min_dims dimensions and no NaN values."""
|
||||
if not isinstance(value, list) or len(value) == 0:
|
||||
return False
|
||||
for item in value:
|
||||
if not isinstance(item, torch.Tensor):
|
||||
return False
|
||||
if item.dim() < min_dims:
|
||||
return False
|
||||
if torch.isnan(item).any().item():
|
||||
return False
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def none_or_tensor_with_dims(dims: int) -> Callable[[Any], bool]:
|
||||
"""Return a validator that checks if value is None or a tensor with specific dimensions and no NaN values."""
|
||||
|
||||
def validator(value: Any) -> bool:
|
||||
if value is None:
|
||||
return True
|
||||
if not isinstance(value, torch.Tensor):
|
||||
return False
|
||||
if value.dim() != dims:
|
||||
return False
|
||||
return not torch.isnan(value).any().item()
|
||||
|
||||
return validator
|
||||
|
||||
@staticmethod
|
||||
def none_or_list(value: Any) -> bool:
|
||||
"""Check if value is None or a list."""
|
||||
return value is None or isinstance(value, list)
|
||||
|
||||
@staticmethod
|
||||
def none_or_positive_int(value: Any) -> bool:
|
||||
"""Check if value is None or a positive integer."""
|
||||
return value is None or (isinstance(value, int) and value > 0)
|
||||
|
||||
# Helper methods that return functions for common patterns
|
||||
@staticmethod
|
||||
def with_dims(dims: int) -> Callable[[Any], bool]:
|
||||
"""Return a validator that checks if tensor has specific dimensions and no NaN values."""
|
||||
|
||||
def validator(value: Any) -> bool:
|
||||
return StageValidators.tensor_with_dims(value, dims)
|
||||
|
||||
return validator
|
||||
|
||||
@staticmethod
|
||||
def min_dims(min_dims: int) -> Callable[[Any], bool]:
|
||||
"""Return a validator that checks if tensor has at least min_dims dimensions and no NaN values."""
|
||||
|
||||
def validator(value: Any) -> bool:
|
||||
return StageValidators.tensor_min_dims(value, min_dims)
|
||||
|
||||
return validator
|
||||
|
||||
@staticmethod
|
||||
def divisible(divisor: int) -> Callable[[Any], bool]:
|
||||
"""Return a validator that checks if value is divisible by divisor."""
|
||||
|
||||
def validator(value: Any) -> bool:
|
||||
return StageValidators.divisible_by(value, divisor)
|
||||
|
||||
return validator
|
||||
|
||||
@staticmethod
|
||||
def positive_int_divisible(divisor: int) -> Callable[[Any], bool]:
|
||||
"""Return a validator that checks if value is a positive integer divisible by divisor."""
|
||||
|
||||
def validator(value: Any) -> bool:
|
||||
return (isinstance(value, int) and value > 0
|
||||
and StageValidators.divisible_by(value, divisor))
|
||||
|
||||
return validator
|
||||
|
||||
@staticmethod
|
||||
def list_of_tensors_dims(dims: int) -> Callable[[Any], bool]:
|
||||
"""Return a validator that checks if value is a list of tensors with specific dimensions and no NaN values."""
|
||||
|
||||
def validator(value: Any) -> bool:
|
||||
return StageValidators.list_of_tensors_with_dims(value, dims)
|
||||
|
||||
return validator
|
||||
|
||||
@staticmethod
|
||||
def list_of_tensors_min_dims(min_dims: int) -> Callable[[Any], bool]:
|
||||
"""Return a validator that checks if value is a list of tensors with at least min_dims dimensions and no NaN values."""
|
||||
|
||||
def validator(value: Any) -> bool:
|
||||
return StageValidators.list_of_tensors_with_min_dims(
|
||||
value, min_dims)
|
||||
|
||||
return validator
|
||||
|
||||
|
||||
class ValidationFailure:
|
||||
"""Details about a specific validation failure."""
|
||||
|
||||
def __init__(self,
|
||||
validator_name: str,
|
||||
actual_value: Any,
|
||||
expected: Optional[str] = None,
|
||||
error_msg: Optional[str] = None):
|
||||
self.validator_name = validator_name
|
||||
self.actual_value = actual_value
|
||||
self.expected = expected
|
||||
self.error_msg = error_msg
|
||||
|
||||
def __str__(self) -> str:
|
||||
parts = [f"Validator '{self.validator_name}' failed"]
|
||||
|
||||
if self.error_msg:
|
||||
parts.append(f"Error: {self.error_msg}")
|
||||
|
||||
# Add actual value info (but limit very long representations)
|
||||
actual_str = self._format_value(self.actual_value)
|
||||
parts.append(f"Actual: {actual_str}")
|
||||
|
||||
if self.expected:
|
||||
parts.append(f"Expected: {self.expected}")
|
||||
|
||||
return ". ".join(parts)
|
||||
|
||||
def _format_value(self, value: Any) -> str:
|
||||
"""Format a value for display in error messages."""
|
||||
if value is None:
|
||||
return "None"
|
||||
elif isinstance(value, torch.Tensor):
|
||||
return f"tensor(shape={list(value.shape)}, dtype={value.dtype})"
|
||||
elif isinstance(value, list):
|
||||
if len(value) == 0:
|
||||
return "[]"
|
||||
elif len(value) <= 3:
|
||||
item_strs = [self._format_value(item) for item in value]
|
||||
return f"[{', '.join(item_strs)}]"
|
||||
else:
|
||||
return f"list(length={len(value)}, first_item={self._format_value(value[0])})"
|
||||
elif isinstance(value, str):
|
||||
if len(value) > 50:
|
||||
return f"'{value[:47]}...'"
|
||||
else:
|
||||
return f"'{value}'"
|
||||
else:
|
||||
return f"{type(value).__name__}({value})"
|
||||
|
||||
|
||||
class VerificationResult:
|
||||
"""Wrapper class for stage verification results."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._checks: Dict[str, bool] = {}
|
||||
self._failures: Dict[str, List[ValidationFailure]] = {}
|
||||
|
||||
def add_check(
|
||||
self, field_name: str, value: Any,
|
||||
validators: Union[Callable[[Any], bool], List[Callable[[Any], bool]]]
|
||||
) -> 'VerificationResult':
|
||||
"""
|
||||
Add a validation check for a field.
|
||||
|
||||
Args:
|
||||
field_name: Name of the field being checked
|
||||
value: The actual value to validate
|
||||
validators: Single validation function or list of validation functions.
|
||||
Each function will be called with the value as its first argument.
|
||||
|
||||
Returns:
|
||||
Self for method chaining
|
||||
|
||||
Examples:
|
||||
# Single validator
|
||||
result.add_check("tensor", my_tensor, V.is_tensor)
|
||||
|
||||
# Multiple validators (all must pass)
|
||||
result.add_check("latents", batch.latents, [V.is_tensor, V.with_dims(5)])
|
||||
|
||||
# Using partial functions for parameters
|
||||
result.add_check("height", batch.height, [V.not_none, V.divisible(8)])
|
||||
"""
|
||||
if not isinstance(validators, list):
|
||||
validators = [validators]
|
||||
|
||||
failures = []
|
||||
all_passed = True
|
||||
|
||||
# Apply all validators and collect detailed failure info
|
||||
for validator in validators:
|
||||
try:
|
||||
passed = validator(value)
|
||||
if not passed:
|
||||
all_passed = False
|
||||
failure = self._create_validation_failure(validator, value)
|
||||
failures.append(failure)
|
||||
except Exception as e:
|
||||
# If any validator raises an exception, consider the check failed
|
||||
all_passed = False
|
||||
validator_name = getattr(validator, '__name__', str(validator))
|
||||
failure = ValidationFailure(
|
||||
validator_name=validator_name,
|
||||
actual_value=value,
|
||||
error_msg=f"Exception during validation: {str(e)}")
|
||||
failures.append(failure)
|
||||
|
||||
self._checks[field_name] = all_passed
|
||||
if not all_passed:
|
||||
self._failures[field_name] = failures
|
||||
|
||||
return self
|
||||
|
||||
def _create_validation_failure(self, validator: Callable,
|
||||
value: Any) -> ValidationFailure:
|
||||
"""Create a ValidationFailure with detailed information."""
|
||||
validator_name = getattr(validator, '__name__', str(validator))
|
||||
|
||||
# Try to extract meaningful expected value info based on validator type
|
||||
expected = None
|
||||
error_msg = None
|
||||
|
||||
# Handle common validator patterns
|
||||
if hasattr(validator, '__closure__') and validator.__closure__:
|
||||
# This is likely a closure (like our helper functions)
|
||||
if 'dims' in validator_name or 'with_dims' in str(validator):
|
||||
if isinstance(value, torch.Tensor):
|
||||
expected = f"tensor with {validator.__closure__[0].cell_contents} dimensions"
|
||||
else:
|
||||
expected = "tensor with specific dimensions"
|
||||
elif 'divisible' in str(validator):
|
||||
expected = f"integer divisible by {validator.__closure__[0].cell_contents}"
|
||||
|
||||
# Handle specific validator types and check for NaN values
|
||||
if validator_name == 'is_tensor':
|
||||
expected = "torch.Tensor without NaN values"
|
||||
if isinstance(value,
|
||||
torch.Tensor) and torch.isnan(value).any().item():
|
||||
error_msg = f"tensor contains {torch.isnan(value).sum().item()} NaN values"
|
||||
elif validator_name == 'positive_int':
|
||||
expected = "positive integer"
|
||||
elif validator_name == 'not_none':
|
||||
expected = "non-None value"
|
||||
elif validator_name == 'list_not_empty':
|
||||
expected = "non-empty list"
|
||||
elif validator_name == 'bool_value':
|
||||
expected = "boolean value"
|
||||
elif 'tensor_with_dims' in validator_name or 'tensor_min_dims' in validator_name:
|
||||
if isinstance(value, torch.Tensor):
|
||||
if torch.isnan(value).any().item():
|
||||
error_msg = f"tensor has {value.dim()} dimensions but contains {torch.isnan(value).sum().item()} NaN values"
|
||||
else:
|
||||
error_msg = f"tensor has {value.dim()} dimensions"
|
||||
elif validator_name == 'is_list':
|
||||
expected = "list"
|
||||
elif validator_name == 'none_or_tensor':
|
||||
expected = "None or tensor without NaN values"
|
||||
if isinstance(value,
|
||||
torch.Tensor) and torch.isnan(value).any().item():
|
||||
error_msg = f"tensor contains {torch.isnan(value).sum().item()} NaN values"
|
||||
elif validator_name == 'list_of_tensors':
|
||||
expected = "non-empty list of tensors without NaN values"
|
||||
if isinstance(value, list) and len(value) > 0:
|
||||
nan_count = 0
|
||||
for item in value:
|
||||
if isinstance(
|
||||
item,
|
||||
torch.Tensor) and torch.isnan(item).any().item():
|
||||
nan_count += torch.isnan(item).sum().item()
|
||||
if nan_count > 0:
|
||||
error_msg = f"list contains tensors with total {nan_count} NaN values"
|
||||
elif 'list_of_tensors_with_dims' in validator_name:
|
||||
expected = "non-empty list of tensors with specific dimensions and no NaN values"
|
||||
if isinstance(value, list) and len(value) > 0:
|
||||
nan_count = 0
|
||||
for item in value:
|
||||
if isinstance(
|
||||
item,
|
||||
torch.Tensor) and torch.isnan(item).any().item():
|
||||
nan_count += torch.isnan(item).sum().item()
|
||||
if nan_count > 0:
|
||||
error_msg = f"list contains tensors with total {nan_count} NaN values"
|
||||
|
||||
return ValidationFailure(validator_name=validator_name,
|
||||
actual_value=value,
|
||||
expected=expected,
|
||||
error_msg=error_msg)
|
||||
|
||||
def is_valid(self) -> bool:
|
||||
"""Check if all validations passed."""
|
||||
return all(self._checks.values())
|
||||
|
||||
def get_failed_fields(self) -> List[str]:
|
||||
"""Get list of fields that failed validation."""
|
||||
return [field for field, passed in self._checks.items() if not passed]
|
||||
|
||||
def get_detailed_failures(self) -> Dict[str, List[ValidationFailure]]:
|
||||
"""Get detailed failure information for each failed field."""
|
||||
return self._failures.copy()
|
||||
|
||||
def get_failure_summary(self) -> str:
|
||||
"""Get a comprehensive summary of all validation failures."""
|
||||
if self.is_valid():
|
||||
return "All validations passed"
|
||||
|
||||
summary_parts = []
|
||||
for field_name, failures in self._failures.items():
|
||||
field_summary = f"\n Field '{field_name}':"
|
||||
for i, failure in enumerate(failures, 1):
|
||||
field_summary += f"\n {i}. {failure}"
|
||||
summary_parts.append(field_summary)
|
||||
|
||||
return "Validation failures:" + "".join(summary_parts)
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
"""Convert to dictionary for backward compatibility."""
|
||||
return self._checks.copy()
|
||||
|
||||
|
||||
# Alias for convenience
|
||||
V = StageValidators
|
||||
@@ -9,7 +9,6 @@ using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
import os
|
||||
from copy import deepcopy
|
||||
from typing import Any, Dict
|
||||
|
||||
import torch
|
||||
@@ -102,21 +101,21 @@ class StepVideoPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
"""
|
||||
Load the modules from the config.
|
||||
"""
|
||||
logger.info("Loading pipeline modules from config: %s", self.config)
|
||||
modules_config = deepcopy(self.config)
|
||||
model_index = self._load_config(self.model_path)
|
||||
logger.info("Loading pipeline modules from config: %s", model_index)
|
||||
|
||||
# remove keys that are not pipeline modules
|
||||
modules_config.pop("_class_name")
|
||||
modules_config.pop("_diffusers_version")
|
||||
model_index.pop("_class_name")
|
||||
model_index.pop("_diffusers_version")
|
||||
|
||||
# some sanity checks
|
||||
assert len(
|
||||
modules_config
|
||||
model_index
|
||||
) > 1, "model_index.json must contain at least one pipeline module"
|
||||
|
||||
required_modules = ["transformer", "scheduler", "vae"]
|
||||
for module_name in required_modules:
|
||||
if module_name not in modules_config:
|
||||
if module_name not in model_index:
|
||||
raise ValueError(
|
||||
f"model_index.json must contain a {module_name} module")
|
||||
logger.info("Diffusers config passed sanity checks")
|
||||
@@ -124,7 +123,7 @@ class StepVideoPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
# all the component models used by the pipeline
|
||||
modules = {}
|
||||
for module_name, (transformers_or_diffusers,
|
||||
architecture) in modules_config.items():
|
||||
architecture) in model_index.items():
|
||||
component_model_path = os.path.join(self.model_path, module_name)
|
||||
module = PipelineComponentLoader.load_module(
|
||||
module_name=module_name,
|
||||
|
||||
@@ -32,7 +32,7 @@ class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.flow_shift)
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
@@ -76,4 +76,36 @@ class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
|
||||
class WanImageToVideoValidationPipeline(ComposedPipelineBase):
|
||||
"""
|
||||
I2V Validation pipeline for Wan2.1, assumes that the input are preprocess latents.
|
||||
"""
|
||||
_required_config_modules = ["vae", "scheduler", "transformer"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer")))
|
||||
|
||||
self.add_stage(stage_name="image_latent_preparation_stage",
|
||||
stage=EncodingStage(vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
|
||||
EntryClass = WanImageToVideoPipeline
|
||||
|
||||
@@ -32,7 +32,7 @@ class WanPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
# We use UniPCMScheduler from Wan2.1 official repo, not the one in diffusers.
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.flow_shift)
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
@@ -75,7 +75,7 @@ class WanValidationPipeline(ComposedPipelineBase):
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.flow_shift)
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
@@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Optional
|
||||
|
||||
from fastvideo.v1.logger import init_logger
|
||||
# imported by other files, do not remove
|
||||
from fastvideo.v1.platforms.interface import _Backend # noqa: F401
|
||||
from fastvideo.v1.platforms.interface import AttentionBackendEnum # noqa: F401
|
||||
from fastvideo.v1.platforms.interface import Platform, PlatformEnum
|
||||
from fastvideo.v1.utils import resolve_obj_by_qualname
|
||||
|
||||
|
||||
@@ -13,8 +13,9 @@ from typing_extensions import ParamSpec
|
||||
|
||||
import fastvideo.v1.envs as envs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.platforms.interface import (DeviceCapability, Platform,
|
||||
PlatformEnum, _Backend)
|
||||
from fastvideo.v1.platforms.interface import (AttentionBackendEnum,
|
||||
DeviceCapability, Platform,
|
||||
PlatformEnum)
|
||||
from fastvideo.v1.utils import import_pynvml
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -106,75 +107,81 @@ class CudaPlatformBase(Platform):
|
||||
return float(torch.cuda.max_memory_allocated(device))
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend: Optional[_Backend],
|
||||
def get_attn_backend_cls(cls,
|
||||
selected_backend: Optional[AttentionBackendEnum],
|
||||
head_size: int, dtype: torch.dtype) -> str:
|
||||
# TODO(will): maybe come up with a more general interface for local attention
|
||||
# if distributed is False, we always try to use Flash attn
|
||||
|
||||
logger.info("Trying FASTVIDEO_ATTENTION_BACKEND=%s",
|
||||
envs.FASTVIDEO_ATTENTION_BACKEND)
|
||||
if selected_backend == _Backend.SLIDING_TILE_ATTN:
|
||||
if selected_backend == AttentionBackendEnum.SLIDING_TILE_ATTN:
|
||||
try:
|
||||
from st_attn import sliding_tile_attention # noqa: F401
|
||||
|
||||
from fastvideo.v1.attention.backends.sliding_tile_attn import ( # noqa: F401
|
||||
SlidingTileAttentionBackend)
|
||||
logger.info("Using Sliding Tile Attention backend.")
|
||||
|
||||
return "fastvideo.v1.attention.backends.sliding_tile_attn.SlidingTileAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
logger.info(
|
||||
"Sliding Tile Attention backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == _Backend.SAGE_ATTN:
|
||||
logger.error(
|
||||
"Failed to import Sliding Tile Attention backend: %s",
|
||||
str(e))
|
||||
raise ImportError(
|
||||
"Sliding Tile Attention backend is not installed. ") from e
|
||||
elif selected_backend == AttentionBackendEnum.SAGE_ATTN:
|
||||
try:
|
||||
from sageattention import sageattn # noqa: F401
|
||||
|
||||
from fastvideo.v1.attention.backends.sage_attn import ( # noqa: F401
|
||||
SageAttentionBackend)
|
||||
logger.info("Using Sage Attention backend.")
|
||||
|
||||
return "fastvideo.v1.attention.backends.sage_attn.SageAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
logger.info(
|
||||
"Sage Attention backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == _Backend.VIDEO_SPARSE_ATTN:
|
||||
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
|
||||
try:
|
||||
from vsa import block_sparse_attn # noqa: F401
|
||||
|
||||
from fastvideo.v1.attention.backends.video_sparse_attn import ( # noqa: F401
|
||||
VideoSparseAttentionBackend)
|
||||
logger.info("Using Video Sparse Attention backend.")
|
||||
|
||||
return "fastvideo.v1.attention.backends.video_sparse_attn.VideoSparseAttentionBackend"
|
||||
except ImportError as e:
|
||||
logger.info(e)
|
||||
logger.info(
|
||||
"Video Sparse Attention backend is not installed. Fall back to Flash Attention."
|
||||
)
|
||||
elif selected_backend == _Backend.TORCH_SDPA:
|
||||
logger.error(
|
||||
"Failed to import Video Sparse Attention backend: %s",
|
||||
str(e))
|
||||
raise ImportError(
|
||||
"Video Sparse Attention backend is not installed. ") from e
|
||||
elif selected_backend == AttentionBackendEnum.TORCH_SDPA:
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
|
||||
elif selected_backend == _Backend.FLASH_ATTN or selected_backend is None:
|
||||
elif selected_backend == AttentionBackendEnum.FLASH_ATTN or selected_backend is None:
|
||||
pass
|
||||
elif selected_backend:
|
||||
raise ValueError(f"Invalid attention backend for {cls.device_name}")
|
||||
|
||||
target_backend = _Backend.FLASH_ATTN
|
||||
target_backend = AttentionBackendEnum.FLASH_ATTN
|
||||
if not cls.has_device_capability(80):
|
||||
logger.info(
|
||||
"Cannot use FlashAttention-2 backend for Volta and Turing "
|
||||
"GPUs.")
|
||||
target_backend = _Backend.TORCH_SDPA
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
elif dtype not in (torch.float16, torch.bfloat16):
|
||||
logger.info(
|
||||
"Cannot use FlashAttention-2 backend for dtype other than "
|
||||
"torch.float16 or torch.bfloat16.")
|
||||
target_backend = _Backend.TORCH_SDPA
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
|
||||
# FlashAttn is valid for the model, checking if the package is
|
||||
# installed.
|
||||
if target_backend == _Backend.FLASH_ATTN:
|
||||
if target_backend == AttentionBackendEnum.FLASH_ATTN:
|
||||
try:
|
||||
import flash_attn # noqa: F401
|
||||
|
||||
@@ -187,19 +194,21 @@ class CudaPlatformBase(Platform):
|
||||
logger.info(
|
||||
"Cannot use FlashAttention-2 backend for head size %d.",
|
||||
head_size)
|
||||
target_backend = _Backend.TORCH_SDPA
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
except ImportError:
|
||||
logger.info("Cannot use FlashAttention-2 backend because the "
|
||||
"flash_attn package is not found. "
|
||||
"Make sure that flash_attn was built and installed "
|
||||
"(on by default).")
|
||||
target_backend = _Backend.TORCH_SDPA
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
|
||||
if target_backend == _Backend.TORCH_SDPA:
|
||||
if target_backend == AttentionBackendEnum.TORCH_SDPA:
|
||||
logger.info("Using Torch SDPA backend.")
|
||||
|
||||
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
|
||||
|
||||
logger.info("Using Flash Attention backend.")
|
||||
|
||||
return "fastvideo.v1.attention.backends.flash_attn.FlashAttentionBackend"
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -13,7 +13,7 @@ from fastvideo.v1.logger import init_logger
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class _Backend(enum.Enum):
|
||||
class AttentionBackendEnum(enum.Enum):
|
||||
FLASH_ATTN = enum.auto()
|
||||
SLIDING_TILE_ATTN = enum.auto()
|
||||
TORCH_SDPA = enum.auto()
|
||||
@@ -88,7 +88,8 @@ class Platform:
|
||||
return self._enum == PlatformEnum.CUDA
|
||||
|
||||
@classmethod
|
||||
def get_attn_backend_cls(cls, selected_backend: Optional[_Backend],
|
||||
def get_attn_backend_cls(cls,
|
||||
selected_backend: Optional[AttentionBackendEnum],
|
||||
head_size: int, dtype: torch.dtype) -> str:
|
||||
"""Get the attention backend class of a device."""
|
||||
return ""
|
||||
@@ -168,6 +169,7 @@ class Platform:
|
||||
random.seed(seed)
|
||||
np.random.seed(seed)
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed)
|
||||
|
||||
@classmethod
|
||||
def verify_model_arch(cls, model_arch: str) -> None:
|
||||
|
||||
@@ -1,8 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import pytest
|
||||
import torch.distributed as dist
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import numpy as np
|
||||
|
||||
|
||||
@@ -10,6 +10,7 @@ from transformers import AutoConfig
|
||||
from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
|
||||
load_tokenizer)
|
||||
# from fastvideo.v1.models.hunyuan.text_encoder import load_text_encoder, load_tokenizer
|
||||
from fastvideo.v1.configs.pipelines import PipelineConfig
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -40,8 +41,7 @@ def test_clip_encoder():
|
||||
- Produce nearly identical outputs for the same input prompts
|
||||
"""
|
||||
args = FastVideoArgs(model_path="openai/clip-vit-large-patch14",
|
||||
text_encoder_precisions=("fp16",),
|
||||
text_encoder_configs=(CLIPTextConfig(),))
|
||||
pipeline_config=PipelineConfig(text_encoder_configs=(CLIPTextConfig(),), text_encoder_precisions=("fp16",)))
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
logger.info("Loading models from %s", args.model_path)
|
||||
|
||||
@@ -8,6 +8,7 @@ from transformers import AutoConfig
|
||||
|
||||
from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
|
||||
load_tokenizer)
|
||||
from fastvideo.v1.configs.pipelines import PipelineConfig
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.logger import init_logger
|
||||
@@ -40,8 +41,7 @@ def test_llama_encoder():
|
||||
- Produce nearly identical outputs for the same input prompts
|
||||
"""
|
||||
args = FastVideoArgs(model_path="meta-llama/Llama-2-7b-hf",
|
||||
text_encoder_precisions=("fp16",),
|
||||
text_encoder_configs=(LlamaConfig(),))
|
||||
pipeline_config=PipelineConfig(text_encoder_configs=(LlamaConfig(),), text_encoder_precisions=("fp16",)))
|
||||
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
|
||||
|
||||
@@ -6,6 +6,7 @@ import pytest
|
||||
import torch
|
||||
from transformers import AutoConfig, AutoTokenizer, UMT5EncoderModel
|
||||
|
||||
from fastvideo.v1.configs.pipelines import PipelineConfig
|
||||
from fastvideo.v1.forward_context import set_forward_context
|
||||
from fastvideo.v1.logger import init_logger
|
||||
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
|
||||
@@ -39,7 +40,8 @@ def test_t5_encoder():
|
||||
precision).to(device).eval()
|
||||
tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
|
||||
|
||||
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH, text_encoder_configs=(T5Config(),), text_encoder_precisions=(precision_str,))
|
||||
|
||||
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH, pipeline_config=PipelineConfig(text_encoder_configs=(T5Config(),), text_encoder_precisions=(precision_str,)))
|
||||
loader = TextEncoderLoader()
|
||||
model2 = loader.load(TEXT_ENCODER_PATH, "", args)
|
||||
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
import os
|
||||
import sys
|
||||
import subprocess
|
||||
from pathlib import Path
|
||||
import pytest
|
||||
|
||||
NUM_NODES = "1"
|
||||
NUM_GPUS_PER_NODE = "1"
|
||||
|
||||
# Set environment variables
|
||||
os.environ["FASTVIDEO_ATTENTION_CONFIG"] = "assets/mask_strategy_wan.json"
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "SLIDING_TILE_ATTN"
|
||||
|
||||
def test_inference():
|
||||
"""Test the inference functionality"""
|
||||
# Create command as in wan_14B-STA.sh
|
||||
cmd = [
|
||||
"fastvideo", "generate",
|
||||
"--model-path", "Wan-AI/Wan2.1-T2V-14B-Diffusers",
|
||||
"--sp-size", "1",
|
||||
"--tp-size", "1",
|
||||
"--num-gpus", "1",
|
||||
"--height", "768",
|
||||
"--width", "1280",
|
||||
"--num-frames", "69",
|
||||
"--num-inference-steps", "2",
|
||||
"--fps", "16",
|
||||
"--guidance-scale", "5.0",
|
||||
"--flow-shift", "5.0",
|
||||
"--prompt", "A majestic lion strides across the golden savanna, its powerful frame glistening under the warm afternoon sun. The tall grass ripples gently in the breeze, enhancing the lion's commanding presence. The tone is vibrant, embodying the raw energy of the wild. Low angle, steady tracking shot, cinematic.",
|
||||
"--negative-prompt", "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards",
|
||||
"--seed", "1024",
|
||||
"--output-path", "outputs_video/STA_1024/",
|
||||
]
|
||||
|
||||
# Run the command
|
||||
subprocess.run(cmd, check=True)
|
||||
|
||||
# Verify output directory exists
|
||||
output_dir = Path("outputs_video/STA_1024/")
|
||||
assert output_dir.exists(), f"Output directory {output_dir} does not exist"
|
||||
|
||||
# Verify that video files were generated
|
||||
video_files = list(output_dir.glob("*.mp4"))
|
||||
assert len(video_files) > 0, "No video files were generated"
|
||||
|
||||
# Verify the video file properties
|
||||
for video_file in video_files:
|
||||
assert video_file.stat().st_size > 0, f"Video file {video_file} is empty"
|
||||
|
||||
if __name__ == "__main__":
|
||||
test_inference()
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user