Compare commits
43
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b118043f55 | ||
|
|
83d45684b6 | ||
|
|
4afb0cfe4f | ||
|
|
3eec1281cf | ||
|
|
0660489e38 | ||
|
|
dd871a17bf | ||
|
|
dc11529862 | ||
|
|
ffabf85e31 | ||
|
|
c0026ca5ba | ||
|
|
afb16377a8 | ||
|
|
b44c28749c | ||
|
|
a9b9bad9a2 | ||
|
|
0f2bbe71ac | ||
|
|
2a46902ecb | ||
|
|
1b2b97544d | ||
|
|
66012d3a4c | ||
|
|
f666b9de41 | ||
|
|
7e3c073b55 | ||
|
|
a6aa21bd07 | ||
|
|
6519b57aab | ||
|
|
675aea6ece | ||
|
|
46e7a15e0d | ||
|
|
e4f702d7ec | ||
|
|
bb68fcc809 | ||
|
|
b392e6a874 | ||
|
|
0991003905 | ||
|
|
8f8ce6d9e1 | ||
|
|
e3d0cbe185 | ||
|
|
d5ec468d43 | ||
|
|
e55fa6e5dc | ||
|
|
61b6ddeee1 | ||
|
|
a9a000f45d | ||
|
|
66b8b8561e | ||
|
|
6684872616 | ||
|
|
7f654e3332 | ||
|
|
8631c1b806 | ||
|
|
5357e12b5a | ||
|
|
bdfdf1dfee | ||
|
|
d156461785 | ||
|
|
6edf113838 | ||
|
|
dcf7738cbc | ||
|
|
b2ebaaf865 | ||
|
|
7768bb80f6 |
@@ -4,14 +4,6 @@ title: "[Bug] "
|
||||
labels: ['Bug']
|
||||
|
||||
body:
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Environment
|
||||
description: |
|
||||
Please share your environment with us. You can run the command **python fastvideo/utils/collect_env.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Describe the bug
|
||||
@@ -25,5 +17,13 @@ body:
|
||||
What command or script did you run? Which **model** are you using?
|
||||
placeholder: |
|
||||
A placeholder for the command.
|
||||
validations:
|
||||
required: true
|
||||
- type: textarea
|
||||
attributes:
|
||||
label: Environment
|
||||
description: |
|
||||
Please share your environment with us. You can run the command **python fastvideo/utils/collect_env.py** and copy-paste its output below.
|
||||
placeholder: FastVideo version, platform, python version, cuda version...
|
||||
validations:
|
||||
required: true
|
||||
@@ -39,6 +39,11 @@ on:
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
run_training_test:
|
||||
description: "Run training-test"
|
||||
required: false
|
||||
default: false
|
||||
type: boolean
|
||||
|
||||
env:
|
||||
PYTHONUNBUFFERED: "1"
|
||||
@@ -59,6 +64,7 @@ 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 }}
|
||||
steps:
|
||||
- uses: actions/checkout@v4
|
||||
- uses: dorny/paths-filter@v3
|
||||
@@ -79,6 +85,8 @@ jobs:
|
||||
- 'fastvideo/v1/tests/transformers/**'
|
||||
- 'fastvideo/v1/layers/**'
|
||||
- 'fastvideo/v1/attention/**'
|
||||
training-test:
|
||||
- 'fastvideo/v1/**'
|
||||
|
||||
encoder-test:
|
||||
needs: change-filter
|
||||
@@ -160,6 +168,25 @@ jobs:
|
||||
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:latest' }}"
|
||||
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/training -srP"
|
||||
timeout_minutes: 30
|
||||
secrets:
|
||||
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
|
||||
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
|
||||
|
||||
runpod-cleanup:
|
||||
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
|
||||
|
||||
@@ -5,7 +5,7 @@ on:
|
||||
branches:
|
||||
- main
|
||||
paths:
|
||||
- "csrc/sliding_tile_attention/setup.py"
|
||||
- "csrc/attn/setup_sta.py"
|
||||
workflow_dispatch:
|
||||
|
||||
jobs:
|
||||
@@ -23,13 +23,13 @@ jobs:
|
||||
- name: Check if version changed
|
||||
id: check-version
|
||||
run: |
|
||||
cd csrc/sliding_tile_attention
|
||||
cd csrc/attn
|
||||
# Get current commit's version
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
|
||||
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_sta.py)
|
||||
echo "New version: $NEW_VERSION"
|
||||
|
||||
# Get previous version from git history
|
||||
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
OLD_VERSION=$(git show HEAD~1:./setup_sta.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
|
||||
echo "Old version: $OLD_VERSION"
|
||||
|
||||
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
|
||||
@@ -136,19 +136,21 @@ jobs:
|
||||
|
||||
- name: Build wheel
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/sliding_tile_attention # Move into the correct folder
|
||||
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py bdist_wheel --dist-dir=dist
|
||||
cd csrc/attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_sta.py bdist_wheel --dist-dir=dist
|
||||
|
||||
- name: Rename wheel file
|
||||
run: |
|
||||
cd csrc/sliding_tile_attention
|
||||
cd csrc/attn
|
||||
|
||||
CUDA_SHORT_VERSION=$(echo ${{ matrix.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
|
||||
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-version }} | cut -d. -f1,2)
|
||||
@@ -163,7 +165,7 @@ jobs:
|
||||
uses: actions/upload-artifact@v4
|
||||
with:
|
||||
name: ${{ env.wheel_name }}
|
||||
path: csrc/sliding_tile_attention/dist/*.whl
|
||||
path: csrc/attn/dist/*.whl
|
||||
retention-days: 90
|
||||
|
||||
publish_package:
|
||||
@@ -229,17 +231,19 @@ jobs:
|
||||
|
||||
- name: Build source distribution
|
||||
run: |
|
||||
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
|
||||
|
||||
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
|
||||
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
|
||||
# However this still fails so I'm using a newer version of setuptools
|
||||
pip install setuptools
|
||||
pip install ninja packaging wheel
|
||||
|
||||
cd csrc/sliding_tile_attention # Move into the correct folder
|
||||
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
|
||||
python setup.py sdist --dist-dir=dist
|
||||
cd csrc/attn # Move into the correct folder
|
||||
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
|
||||
python setup_sta.py sdist --dist-dir=dist
|
||||
|
||||
- name: Publish release distributions to PyPI
|
||||
uses: pypa/gh-action-pypi-publish@release/v1
|
||||
with:
|
||||
packages-dir: csrc/sliding_tile_attention/dist/
|
||||
packages-dir: csrc/attn/dist/
|
||||
|
||||
@@ -28,4 +28,4 @@ jobs:
|
||||
|
||||
- name: Run Pytest
|
||||
run: |
|
||||
pytest --ignore csrc/sliding_tile_attention/test
|
||||
pytest --ignore csrc/attn/test
|
||||
|
||||
@@ -27,7 +27,6 @@ env
|
||||
**/build/
|
||||
**.pyc
|
||||
**.txt
|
||||
**.json
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
|
||||
+2
-2
@@ -1,3 +1,3 @@
|
||||
[submodule "csrc/sliding_tile_attention/tk"]
|
||||
path = csrc/sliding_tile_attention/tk
|
||||
[submodule "csrc/attn/tk"]
|
||||
path = csrc/attn/tk
|
||||
url = https://github.com/HazyResearch/ThunderKittens.git
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -6,7 +6,6 @@
|
||||
## 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
|
||||
@@ -16,17 +15,27 @@ sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave
|
||||
sudo apt update
|
||||
sudo apt install clang-11
|
||||
```
|
||||
Install STA:
|
||||
|
||||
## Environment Setup
|
||||
First, set up your CUDA environment:
|
||||
```bash
|
||||
export CUDA_HOME=/usr/local/cuda-12.4
|
||||
export PATH=${CUDA_HOME}/bin:${PATH}
|
||||
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
|
||||
git submodule update --init --recursive
|
||||
python setup.py install
|
||||
```
|
||||
|
||||
## Install Sliding Tile Attention (STA)
|
||||
```bash
|
||||
python setup_sta.py install
|
||||
```
|
||||
|
||||
## Install Video Sparse Attention (VSA)
|
||||
```bash
|
||||
python setup_vsa.py install
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
from st_attn import sliding_tile_attention
|
||||
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
|
||||
@@ -45,12 +45,12 @@ def benchmark_attention(configurations):
|
||||
|
||||
# Warmup for forward pass
|
||||
for _ in range(10):
|
||||
o = sliding_tile_attention(q, k, v, [[6, 6, 6]] * 24, 0, False)
|
||||
o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, '18x48x80')
|
||||
|
||||
# Time the forward pass
|
||||
for i in range(10):
|
||||
start_events_fwd[i].record()
|
||||
o = sliding_tile_attention(q, k, v, [[6, 6, 6]] * 24, 0, False)
|
||||
o = sliding_tile_attention(q, k, v, [[3, 6, 10]] * 24, 0, False, '18x48x80')
|
||||
end_events_fwd[i].record()
|
||||
|
||||
torch.cuda.synchronize()
|
||||
@@ -124,7 +124,7 @@ def plot_results(results):
|
||||
|
||||
# Example list of configurations to test
|
||||
configurations = [
|
||||
(2, 24, 82944, 128, False),
|
||||
(2, 24, 69120, 128, False),
|
||||
# (16, 16, 768*16, 128, False),
|
||||
# (16, 16, 768*2, 128, False),
|
||||
# (16, 16, 768*4, 128, False),
|
||||
@@ -0,0 +1,225 @@
|
||||
import torch
|
||||
import argparse
|
||||
from flash_attn.utils.benchmark import benchmark_forward
|
||||
from vsa import block_sparse_attention_fwd, block_sparse_attention_backward
|
||||
from vsa import BLOCK_M, BLOCK_N
|
||||
|
||||
import numpy as np
|
||||
import random
|
||||
|
||||
def set_seed(seed: int = 42):
|
||||
# Python random module
|
||||
random.seed(seed)
|
||||
|
||||
# NumPy
|
||||
np.random.seed(seed)
|
||||
|
||||
# PyTorch
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed) # if using multi-GPU
|
||||
|
||||
def parse_arguments():
|
||||
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
|
||||
parser.add_argument('--batch_size', type=int, default=1, help='Batch size')
|
||||
parser.add_argument('--num_heads', type=int, default=12, help='Number of heads')
|
||||
parser.add_argument('--head_dim', type=int, default=64, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=None, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[49152], help='Sequence lengths to benchmark')
|
||||
return parser.parse_args()
|
||||
|
||||
def create_input_tensors(batch, head, seq_len, headdim):
|
||||
"""Create random input tensors for attention."""
|
||||
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
return q, k, v
|
||||
|
||||
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
|
||||
|
||||
Args:
|
||||
bs: batch size
|
||||
h: number of heads
|
||||
num_q_blocks: number of query blocks
|
||||
num_kv_blocks: number of key-value blocks
|
||||
k: number of kv blocks each q block attends to
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to k).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Binary mask where 1 indicates attention connection.
|
||||
"""
|
||||
# Ensure k is not larger than num_kv_blocks
|
||||
k = min(k, num_kv_blocks)
|
||||
|
||||
# Create random scores for sampling
|
||||
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
|
||||
|
||||
# Get top-k indices for each q block
|
||||
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
|
||||
|
||||
# sort q2k_block_sparse_index
|
||||
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
|
||||
|
||||
# All q blocks attend to exactly k kv blocks
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
|
||||
|
||||
# Create the corresponding mask
|
||||
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
|
||||
|
||||
# Fill in the mask based on the indices
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx]
|
||||
block_sparse_mask[b, head, q_idx, kv_indices] = True
|
||||
|
||||
# Create the reverse mapping (k2q)
|
||||
# First, initialize lists to collect q indices for each kv block
|
||||
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
|
||||
|
||||
# Populate the lists based on q2k mapping
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
|
||||
for kv_idx in kv_indices:
|
||||
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
|
||||
|
||||
# Find the maximum number of q blocks that attend to any kv block
|
||||
max_q_per_kv = 0
|
||||
for flat_idx in range(bs * h):
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
|
||||
|
||||
# Create tensors for k2q mapping
|
||||
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
|
||||
dtype=torch.int32, device=device)
|
||||
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
|
||||
dtype=torch.int32, device=device)
|
||||
|
||||
# Fill the tensors
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
q_indices = k2q_indices_list[flat_idx][kv_idx]
|
||||
num_q = len(q_indices)
|
||||
k2q_block_sparse_num[b, head, kv_idx] = num_q
|
||||
if num_q > 0:
|
||||
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
|
||||
q_indices, dtype=torch.int32, device=device)
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
|
||||
|
||||
def benchmark_block_sparse_attention(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops):
|
||||
"""Benchmark block sparse attention forward and backward passes."""
|
||||
print("\n=== BLOCK SPARSE ATTENTION BENCHMARK ===")
|
||||
|
||||
# Forward pass
|
||||
# Warm-up run
|
||||
o, l_vec = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Benchmark forward
|
||||
_, fwd_time = benchmark_forward(
|
||||
block_sparse_attention_fwd,
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num,
|
||||
repeats=20,
|
||||
verbose=False,
|
||||
desc='Block Sparse Forward'
|
||||
)
|
||||
|
||||
sparse_tflops = flops / fwd_time.mean * 1e-12
|
||||
print(f"Block Sparse Forward - TFLOPS: {sparse_tflops:.2f}")
|
||||
|
||||
# Backward pass
|
||||
grad_output = torch.randn_like(o)
|
||||
|
||||
# Warm-up runs
|
||||
for _ in range(5):
|
||||
block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
torch.cuda.synchronize()
|
||||
|
||||
# Benchmark backward
|
||||
_, bwd_time = benchmark_forward(
|
||||
block_sparse_attention_backward,
|
||||
q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num,
|
||||
repeats=20,
|
||||
verbose=False,
|
||||
desc='Block Sparse Backward'
|
||||
)
|
||||
bwd_flops = 2.5 * flops # Approximation
|
||||
|
||||
sparse_bwd_tflops = bwd_flops / bwd_time.mean * 1e-12
|
||||
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd_tflops:.2f}")
|
||||
|
||||
return sparse_tflops, sparse_bwd_tflops
|
||||
|
||||
def main():
|
||||
args = parse_arguments()
|
||||
|
||||
set_seed(42)
|
||||
|
||||
# Extract parameters
|
||||
batch = args.batch_size
|
||||
head = args.num_heads
|
||||
headdim = args.head_dim
|
||||
|
||||
print(f"Block Sparse Attention Benchmark")
|
||||
print(f"batch: {batch}, head: {head}, headdim: {headdim}")
|
||||
|
||||
# Test with different sequence lengths
|
||||
for seq_len in args.seq_lengths:
|
||||
# Skip very long sequences if they might cause OOM
|
||||
if seq_len > 16384 and batch > 1:
|
||||
continue
|
||||
|
||||
print("="*100)
|
||||
print(f"\nSequence length: {seq_len}")
|
||||
|
||||
# Calculate theoretical FLOPs for attention
|
||||
flops = 4 * batch * head * headdim * seq_len * seq_len
|
||||
|
||||
# Create input tensors
|
||||
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
|
||||
|
||||
# Setup block sparse parameters
|
||||
num_q_blocks = seq_len // BLOCK_M
|
||||
num_kv_blocks = seq_len // BLOCK_N
|
||||
|
||||
# Determine k value (number of kv blocks per q block)
|
||||
topk = args.topk
|
||||
if topk is None:
|
||||
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
|
||||
topk = max(1, topk)
|
||||
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
|
||||
|
||||
# Generate block sparse pattern
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, _ = generate_block_sparse_pattern(
|
||||
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
|
||||
|
||||
# Benchmark block sparse attention
|
||||
sparse_fwd, sparse_bwd = benchmark_block_sparse_attention(
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, flops
|
||||
)
|
||||
|
||||
# Print results
|
||||
print("\n=== PERFORMANCE RESULTS ===")
|
||||
print(f"Block Sparse Forward - TFLOPS: {sparse_fwd:.2f}")
|
||||
print(f"Block Sparse Backward - TFLOPS: {sparse_bwd:.2f}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,6 +1,6 @@
|
||||
### ADD TO THIS TO REGISTER NEW KERNELS
|
||||
sources = {
|
||||
'attn': {
|
||||
'st_attn': {
|
||||
'source_files': {
|
||||
'h100': 'st_attn/st_attn_h100.cu' # define these source files for each GPU target desired.
|
||||
}
|
||||
@@ -9,7 +9,7 @@ sources = {
|
||||
|
||||
### WHICH KERNELS DO WE WANT TO BUILD?
|
||||
# (oftentimes during development work you don't need to redefine them all.)
|
||||
kernels = ['attn']
|
||||
kernels = ['st_attn']
|
||||
|
||||
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
|
||||
target = 'h100'
|
||||
@@ -0,0 +1,15 @@
|
||||
### ADD TO THIS TO REGISTER NEW KERNELS
|
||||
sources = {
|
||||
'block_sparse': {
|
||||
'source_files': {
|
||||
'h100': 'vsa/block_sparse_h100.cu'
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
### WHICH KERNELS DO WE WANT TO BUILD?
|
||||
# (oftentimes during development work you don't need to redefine them all.)
|
||||
kernels = ['block_sparse']
|
||||
|
||||
### WHICH GPU TARGET DO WE WANT TO BUILD FOR?
|
||||
target = 'h100'
|
||||
@@ -1,7 +1,7 @@
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
from config import kernels, sources, target
|
||||
from csrc.attn.config_sta import kernels, sources, target
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
@@ -0,0 +1,76 @@
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
from csrc.attn.config_vsa import kernels, sources, target
|
||||
from setuptools import find_packages, setup
|
||||
from torch.utils.cpp_extension import BuildExtension, CUDAExtension
|
||||
|
||||
target = target.lower()
|
||||
|
||||
# Package metadata
|
||||
PACKAGE_NAME = "vsa"
|
||||
VERSION = "0.0.1"
|
||||
AUTHOR = "Hao AI Lab"
|
||||
DESCRIPTION = "Video Sparse Attention Kernel Used in FastVideo"
|
||||
URL = "https://github.com/hao-ai-lab/FastVideo/tree/main/csrc/attn"
|
||||
|
||||
# Set environment variables
|
||||
tk_root = os.getenv('THUNDERKITTENS_ROOT', os.path.abspath(os.path.join(os.getcwd(), 'tk/')))
|
||||
python_include = subprocess.check_output(['python', '-c',
|
||||
"import sysconfig; print(sysconfig.get_path('include'))"]).decode().strip()
|
||||
torch_include = subprocess.check_output([
|
||||
'python', '-c',
|
||||
"import torch; from torch.utils.cpp_extension import include_paths; print(' '.join(['-I' + p for p in include_paths()]))"
|
||||
]).decode().strip()
|
||||
print('vsa root:', tk_root)
|
||||
print('Python include:', python_include)
|
||||
print('Torch include directories:', torch_include)
|
||||
|
||||
# CUDA flags
|
||||
cuda_flags = [
|
||||
'-DNDEBUG', '-Xcompiler=-Wno-psabi', '-Xcompiler=-fno-strict-aliasing', '--expt-extended-lambda',
|
||||
'--expt-relaxed-constexpr', '-forward-unknown-to-host-compiler', '--use_fast_math', '-std=c++20', '-O3',
|
||||
'-Xnvlink=--verbose', '-Xptxas=--verbose', '-Xptxas=--warn-on-spills', f'-I{tk_root}/include',
|
||||
f'-I{tk_root}/prototype', f'-I{python_include}', '-DTORCH_COMPILE'
|
||||
] + torch_include.split()
|
||||
cpp_flags = ['-std=c++20', '-O3']
|
||||
|
||||
if target == 'h100':
|
||||
cuda_flags.append('-DKITTENS_HOPPER')
|
||||
cuda_flags.append('-arch=sm_90a')
|
||||
else:
|
||||
raise ValueError(f'Target {target} not supported')
|
||||
|
||||
source_files = ['vsa.cpp']
|
||||
for k in kernels:
|
||||
if target not in sources[k]['source_files']:
|
||||
raise KeyError(f'Target {target} not found in source files for kernel {k}')
|
||||
if isinstance(sources[k]['source_files'][target], list):
|
||||
source_files.extend(sources[k]['source_files'][target])
|
||||
else:
|
||||
source_files.append(sources[k]['source_files'][target])
|
||||
cpp_flags.append(f'-DTK_COMPILE_{k.replace(" ", "_").upper()}')
|
||||
|
||||
setup(name=PACKAGE_NAME,
|
||||
version=VERSION,
|
||||
author=AUTHOR,
|
||||
description=DESCRIPTION,
|
||||
url=URL,
|
||||
packages=find_packages(),
|
||||
ext_modules=[
|
||||
CUDAExtension('vsa_cuda',
|
||||
sources=source_files,
|
||||
extra_compile_args={
|
||||
'cxx': cpp_flags,
|
||||
'nvcc': cuda_flags
|
||||
},
|
||||
libraries=['cuda'])
|
||||
],
|
||||
cmdclass={'build_ext': BuildExtension},
|
||||
classifiers=[
|
||||
"Programming Language :: Python :: 3",
|
||||
"Environment :: GPU :: NVIDIA CUDA :: 12",
|
||||
"License :: OSI Approved :: Apache Software License",
|
||||
],
|
||||
python_requires='>=3.10',
|
||||
install_requires=["torch>=2.5.0"])
|
||||
@@ -7,8 +7,7 @@
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
|
||||
#ifdef TK_COMPILE_ATTN
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
extern torch::Tensor sta_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, int kernel_t_size, int kernel_w_size, int kernel_h_size, int text_length, bool process_text, bool has_text, int kernel_aspect_ratio_flag
|
||||
);
|
||||
@@ -17,8 +16,8 @@ extern torch::Tensor sta_forward(
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "Sliding Block Attention Kernels"; // optional module docstring
|
||||
|
||||
#ifdef TK_COMPILE_ATTN
|
||||
|
||||
#ifdef TK_COMPILE_ST_ATTN
|
||||
m.def("sta_fwd", torch::wrap_pybind_function(sta_forward), "sliding tile attention, assuming tile size is (6,8,8)");
|
||||
#endif
|
||||
}
|
||||
}
|
||||
@@ -1,19 +1,22 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
from st_attn_cuda import sta_fwd
|
||||
from torch.utils.checkpoint import detach_variable
|
||||
try:
|
||||
from st_attn_cuda import sta_fwd
|
||||
except ImportError:
|
||||
sta_fwd = None
|
||||
|
||||
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, img_latent_shape='30*48*80'):
|
||||
def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_text=True, dit_seq_shape='30x48x80'):
|
||||
seq_length = q_all.shape[2]
|
||||
img_latent_shape_mapping = {
|
||||
dit_seq_shape_mapping = {
|
||||
'30x48x80':1,
|
||||
'36x48x48':2,
|
||||
'18x48x80':3,
|
||||
}
|
||||
if has_text:
|
||||
assert q_all.shape[
|
||||
2] >= 115200, "STA currently only supports video with latent size (30, 48, 80), which is 117 frames x 768 x 1280 pixels"
|
||||
2] >= 115200 and q_all.shape[2] <= 115456, f"Unsupported {dit_seq_shape}, current shape is {q_all.shape}, only support '30x48x80' for HunyuanVideo"
|
||||
assert q_all.shape[1] == len(window_size), "Number of heads must match the number of window sizes"
|
||||
target_size = math.ceil(seq_length / 384) * 384
|
||||
pad_size = target_size - seq_length
|
||||
@@ -22,14 +25,14 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
|
||||
k_all = torch.cat([k_all, k_all[:, :, -pad_size:]], dim=2)
|
||||
v_all = torch.cat([v_all, v_all[:, :, -pad_size:]], dim=2)
|
||||
else:
|
||||
if img_latent_shape == '36x48x48': # Stepvideo 204x768x68
|
||||
if dit_seq_shape == '36x48x48': # Stepvideo 204x768x68
|
||||
assert q_all.shape[2] == 82944
|
||||
elif img_latent_shape == '18x48x80': # Wan 69x768x1280
|
||||
elif dit_seq_shape == '18x48x80': # Wan 69x768x1280
|
||||
assert q_all.shape[2] == 69120
|
||||
else:
|
||||
raise ValueError(f"Unsupported {img_latent_shape}, current shape is {q_all.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
|
||||
raise ValueError(f"Unsupported {dit_seq_shape}, current shape is {q_all.shape}, only support '36x48x48' for Stepvideo and '18x48x80' for Wan")
|
||||
|
||||
kernel_aspect_ratio_flag = img_latent_shape_mapping[img_latent_shape]
|
||||
kernel_aspect_ratio_flag = dit_seq_shape_mapping[dit_seq_shape]
|
||||
hidden_states = torch.empty_like(q_all)
|
||||
# This for loop is ugly. but it is actually quite efficient. The sequence dimension alone can already oversubscribe SMs
|
||||
for head_index, (t_kernel, h_kernel, w_kernel) in enumerate(window_size):
|
||||
@@ -43,4 +46,4 @@ def sliding_tile_attention(q_all, k_all, v_all, window_size, text_length, has_te
|
||||
_ = sta_fwd(q_head, k_head, v_head, o_head, t_kernel, h_kernel, w_kernel, text_length, False, has_text, kernel_aspect_ratio_flag)
|
||||
if has_text:
|
||||
_ = sta_fwd(q_all, k_all, v_all, hidden_states, 3, 3, 3, text_length, True, True, kernel_aspect_ratio_flag)
|
||||
return hidden_states[:, :, :seq_length]
|
||||
return hidden_states[:, :, :seq_length]
|
||||
+1
@@ -829,3 +829,4 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
|
||||
return o;
|
||||
cudaDeviceSynchronize();
|
||||
}
|
||||
|
||||
@@ -0,0 +1,266 @@
|
||||
import torch
|
||||
import argparse
|
||||
from flash_attn.utils.benchmark import benchmark_forward
|
||||
from flash_attn import flash_attn_func
|
||||
from vsa import block_sparse_attention_fwd, block_sparse_attention_backward, BlockSparseAttentionFunction
|
||||
from vsa import BLOCK_M, BLOCK_N
|
||||
|
||||
import numpy as np
|
||||
import random
|
||||
|
||||
def set_seed(seed: int = 42):
|
||||
# Python random module
|
||||
random.seed(seed)
|
||||
|
||||
# NumPy
|
||||
np.random.seed(seed)
|
||||
|
||||
# PyTorch
|
||||
torch.manual_seed(seed)
|
||||
torch.cuda.manual_seed(seed)
|
||||
torch.cuda.manual_seed_all(seed) # if using multi-GPU
|
||||
|
||||
def parse_arguments():
|
||||
parser = argparse.ArgumentParser(description='Benchmark Block Sparse Attention')
|
||||
parser.add_argument('--batch_size', type=int, default=4, help='Batch size')
|
||||
parser.add_argument('--num_heads', type=int, default=6, help='Number of heads')
|
||||
parser.add_argument('--head_dim', type=int, default=128, help='Head dimension')
|
||||
parser.add_argument('--topk', type=int, default=64, help='Number of kv blocks each q block attends to')
|
||||
parser.add_argument('--seq_lengths', type=int, nargs='+', default=[29120], help='Sequence lengths to benchmark')
|
||||
parser.add_argument('--num_iterations', type=int, default=100, help='Number of test iterations to run')
|
||||
return parser.parse_args()
|
||||
|
||||
@torch.no_grad
|
||||
def precision_metric(quant_o, fa2_o):
|
||||
x, xx = quant_o.float(), fa2_o.float()
|
||||
sim = torch.nn.functional.cosine_similarity(x.reshape(1, -1), xx.reshape(1, -1)).item()
|
||||
l1 = ((x - xx).abs().sum() / xx.abs().sum() ).item()
|
||||
rmse = torch.sqrt(torch.mean((x -xx) ** 2)).item()
|
||||
|
||||
return sim, l1, rmse
|
||||
|
||||
def create_input_tensors(batch, head, seq_len, headdim):
|
||||
"""Create random input tensors for attention."""
|
||||
q = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
k = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
v = torch.randn(batch, head, seq_len, headdim, dtype=torch.bfloat16, device="cuda")
|
||||
return q, k, v
|
||||
|
||||
def generate_block_sparse_pattern(bs, h, num_q_blocks, num_kv_blocks, k, device="cuda"):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly k kv blocks.
|
||||
|
||||
Args:
|
||||
bs: batch size
|
||||
h: number of heads
|
||||
num_q_blocks: number of query blocks
|
||||
num_kv_blocks: number of key-value blocks
|
||||
k: number of kv blocks each q block attends to
|
||||
device: device to create tensors on
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, k]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to k).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, num_q_blocks]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Binary mask where 1 indicates attention connection.
|
||||
"""
|
||||
# Ensure k is not larger than num_kv_blocks
|
||||
k = min(k, num_kv_blocks)
|
||||
|
||||
# Create random scores for sampling
|
||||
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device=device)
|
||||
|
||||
# Get top-k indices for each q block
|
||||
_, q2k_block_sparse_index = torch.topk(scores, k, dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(torch.int32)
|
||||
|
||||
# sort q2k_block_sparse_index
|
||||
q2k_block_sparse_index, _ = torch.sort(q2k_block_sparse_index, dim=-1)
|
||||
|
||||
# All q blocks attend to exactly k kv blocks
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks), k, dtype=torch.int32, device=device)
|
||||
|
||||
# Create the corresponding mask
|
||||
block_sparse_mask = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
|
||||
|
||||
# Fill in the mask based on the indices
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx]
|
||||
block_sparse_mask[b, head, q_idx, kv_indices] = True
|
||||
|
||||
# Create the reverse mapping (k2q)
|
||||
# First, initialize lists to collect q indices for each kv block
|
||||
k2q_indices_list = [[[] for _ in range(num_kv_blocks)] for _ in range(bs * h)]
|
||||
|
||||
# Populate the lists based on q2k mapping
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for q_idx in range(num_q_blocks):
|
||||
kv_indices = q2k_block_sparse_index[b, head, q_idx].tolist()
|
||||
for kv_idx in kv_indices:
|
||||
k2q_indices_list[flat_idx][kv_idx].append(q_idx)
|
||||
|
||||
# Find the maximum number of q blocks that attend to any kv block
|
||||
max_q_per_kv = 0
|
||||
for flat_idx in range(bs * h):
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
max_q_per_kv = max(max_q_per_kv, len(k2q_indices_list[flat_idx][kv_idx]))
|
||||
|
||||
# Create tensors for k2q mapping
|
||||
k2q_block_sparse_index = torch.full((bs, h, num_kv_blocks, max_q_per_kv), -1,
|
||||
dtype=torch.int32, device=device)
|
||||
k2q_block_sparse_num = torch.zeros((bs, h, num_kv_blocks),
|
||||
dtype=torch.int32, device=device)
|
||||
|
||||
# Fill the tensors
|
||||
for b in range(bs):
|
||||
for head in range(h):
|
||||
flat_idx = b * h + head
|
||||
for kv_idx in range(num_kv_blocks):
|
||||
q_indices = k2q_indices_list[flat_idx][kv_idx]
|
||||
num_q = len(q_indices)
|
||||
k2q_block_sparse_num[b, head, kv_idx] = num_q
|
||||
if num_q > 0:
|
||||
k2q_block_sparse_index[b, head, kv_idx, :num_q] = torch.tensor(
|
||||
q_indices, dtype=torch.int32, device=device)
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask
|
||||
|
||||
def main():
|
||||
args = parse_arguments()
|
||||
|
||||
set_seed(42)
|
||||
|
||||
# Extract parameters
|
||||
batch = args.batch_size
|
||||
head = args.num_heads
|
||||
headdim = args.head_dim
|
||||
num_iterations = args.num_iterations
|
||||
|
||||
print(f"Block Sparse Attention Benchmark")
|
||||
print(f"batch: {batch}, head: {head}, headdim: {headdim}, iterations: {num_iterations}")
|
||||
|
||||
# Test with different sequence lengths
|
||||
for seq_len in args.seq_lengths:
|
||||
# Skip very long sequences if they might cause OOM
|
||||
# if seq_len > 16384 and batch > 1:
|
||||
# continue
|
||||
|
||||
print("="*100)
|
||||
print(f"\nSequence length: {seq_len}")
|
||||
|
||||
# Collect metrics across iterations
|
||||
forward_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_q_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_k_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
grad_v_metrics = {'sim': [], 'l1': [], 'rmse': []}
|
||||
|
||||
for iter_idx in range(num_iterations):
|
||||
if num_iterations > 1:
|
||||
print(f"\nIteration {iter_idx+1}/{num_iterations}")
|
||||
|
||||
# Create input tensors
|
||||
q, k, v = create_input_tensors(batch, head, seq_len, headdim)
|
||||
|
||||
# Setup block sparse parameters
|
||||
num_q_blocks = seq_len // BLOCK_M
|
||||
num_kv_blocks = seq_len // BLOCK_N
|
||||
|
||||
# Determine k value (number of kv blocks per q block)
|
||||
topk = args.topk
|
||||
if topk is None:
|
||||
topk = num_kv_blocks // 10 # Default to ~90% sparsity if k is not specified
|
||||
topk = max(1, topk)
|
||||
if iter_idx == 0: # Only print this once
|
||||
print(f"Using topk={topk} kv blocks per q block (out of {num_kv_blocks} total kv blocks)")
|
||||
|
||||
# Generate block sparse pattern
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num, block_sparse_mask = generate_block_sparse_pattern(
|
||||
batch, head, num_q_blocks, num_kv_blocks, topk, device="cuda")
|
||||
|
||||
# expand block_sparse_mask to full mask
|
||||
block_mask_expanded = block_sparse_mask.unsqueeze(-1).unsqueeze(-2) # [b, h, num_q_blocks, num_kv_blocks, 1, 1]
|
||||
block_mask_expanded = block_mask_expanded.expand(-1, -1, -1, -1, BLOCK_M, BLOCK_N) # [b, h, num_q_blocks, num_kv_blocks, BLOCK_M, BLOCK_N]
|
||||
full_mask = block_mask_expanded.permute(0, 1, 2, 4, 3, 5).reshape(batch, head, seq_len, seq_len)
|
||||
|
||||
q_sdpa = q.clone()
|
||||
k_sdpa = k.clone()
|
||||
v_sdpa = v.clone()
|
||||
|
||||
q.requires_grad = True
|
||||
k.requires_grad = True
|
||||
v.requires_grad = True
|
||||
q_sdpa.requires_grad = True
|
||||
k_sdpa.requires_grad = True
|
||||
v_sdpa.requires_grad = True
|
||||
|
||||
# testing forward
|
||||
o = BlockSparseAttentionFunction.apply(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
o_sdpa = torch.nn.functional.scaled_dot_product_attention(q_sdpa, k_sdpa, v_sdpa, attn_mask=full_mask)
|
||||
|
||||
sim, l1, rmse = precision_metric(o, o_sdpa)
|
||||
forward_metrics['sim'].append(sim)
|
||||
forward_metrics['l1'].append(l1)
|
||||
forward_metrics['rmse'].append(rmse)
|
||||
|
||||
print(f"block_sparse_attention_fwd vs torch.nn.functional.scaled_dot_product_attention:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
# test backward
|
||||
grad_o = torch.randn_like(o)
|
||||
o.backward(grad_o)
|
||||
o_sdpa.backward(grad_o)
|
||||
|
||||
sim, l1, rmse = precision_metric(q.grad, q_sdpa.grad)
|
||||
grad_q_metrics['sim'].append(sim)
|
||||
grad_q_metrics['l1'].append(l1)
|
||||
grad_q_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_q:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
sim, l1, rmse = precision_metric(k.grad, k_sdpa.grad)
|
||||
grad_k_metrics['sim'].append(sim)
|
||||
grad_k_metrics['l1'].append(l1)
|
||||
grad_k_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_k:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
sim, l1, rmse = precision_metric(v.grad, v_sdpa.grad)
|
||||
grad_v_metrics['sim'].append(sim)
|
||||
grad_v_metrics['l1'].append(l1)
|
||||
grad_v_metrics['rmse'].append(rmse)
|
||||
print(f"block_sparse_attention_bwd vs torch.nn.functional.scaled_dot_product_attention grad_v:\nsim: {sim}, l1: {l1}, rmse: {rmse}")
|
||||
|
||||
# Print summary statistics if multiple iterations were run
|
||||
if num_iterations > 1:
|
||||
print("\n" + "="*50)
|
||||
print(f"Summary Statistics (over {num_iterations} iterations):")
|
||||
|
||||
print("\nForward metrics:")
|
||||
print(f"Similarity: mean={np.mean(forward_metrics['sim']):.6f}, std={np.std(forward_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(forward_metrics['l1']):.6f}, std={np.std(forward_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(forward_metrics['rmse']):.6f}, std={np.std(forward_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient Q metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_q_metrics['sim']):.6f}, std={np.std(grad_q_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_q_metrics['l1']):.6f}, std={np.std(grad_q_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_q_metrics['rmse']):.6f}, std={np.std(grad_q_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient K metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_k_metrics['sim']):.6f}, std={np.std(grad_k_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_k_metrics['l1']):.6f}, std={np.std(grad_k_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_k_metrics['rmse']):.6f}, std={np.std(grad_k_metrics['rmse']):.6f}")
|
||||
|
||||
print("\nGradient V metrics:")
|
||||
print(f"Similarity: mean={np.mean(grad_v_metrics['sim']):.6f}, std={np.std(grad_v_metrics['sim']):.6f}")
|
||||
print(f"L1 error: mean={np.mean(grad_v_metrics['l1']):.6f}, std={np.std(grad_v_metrics['l1']):.6f}")
|
||||
print(f"RMSE: mean={np.mean(grad_v_metrics['rmse']):.6f}, std={np.std(grad_v_metrics['rmse']):.6f}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,136 @@
|
||||
import torch
|
||||
from tqdm import tqdm
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
|
||||
def pytorch_test(Q, K, V, dO):
|
||||
q_ = Q.to(torch.float64).requires_grad_()
|
||||
k_ = K.to(torch.float64).requires_grad_()
|
||||
v_ = V.to(torch.float64).requires_grad_()
|
||||
dO_ = dO.to(torch.float64)
|
||||
|
||||
# manual pytorch implementation of scaled dot product attention
|
||||
QK = torch.matmul(q_, k_.transpose(-2, -1))
|
||||
QK /= (q_.size(-1) ** 0.5)
|
||||
|
||||
# Causal mask removed since causal is always false
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v_)
|
||||
|
||||
output.backward(dO_)
|
||||
|
||||
q_grad = q_.grad
|
||||
k_grad = k_.grad
|
||||
v_grad = v_.grad
|
||||
|
||||
return output, q_grad, k_grad, v_grad
|
||||
|
||||
def fa2_test(Q, K, V, dO):
|
||||
Q.requires_grad = True
|
||||
K.requires_grad = True
|
||||
V.requires_grad = True
|
||||
output = torch.nn.functional.scaled_dot_product_attention(Q, K, V, is_causal=False)
|
||||
output.backward(dO)
|
||||
|
||||
return output, Q.grad, K.grad, V.grad
|
||||
|
||||
def generate_tensor(shape, mean, std, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
|
||||
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
|
||||
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
|
||||
|
||||
return scaled_tensor.contiguous()
|
||||
|
||||
def check_correctness(b, h, n, d, mean, std, num_iterations=100, error_mode='all', test_mode='forward_backward'):
|
||||
results = {
|
||||
'FA2 vs PT': {'sum_diff': 0, 'sum_abs': 0, 'max_diff': 0},
|
||||
}
|
||||
|
||||
for _ in range(num_iterations):
|
||||
torch.manual_seed(0)
|
||||
|
||||
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
dO = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
|
||||
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, dO)
|
||||
fa2_o, fa2_qg, fa2_kg, fa2_vg = fa2_test(Q, K, V, dO)
|
||||
|
||||
if test_mode == 'forward_only':
|
||||
tensors_fa2_pt = [(pt_o, fa2_o)]
|
||||
else: # 'forward_backward'
|
||||
if error_mode == 'output':
|
||||
tensors_fa2_pt = [(pt_o, fa2_o)]
|
||||
elif error_mode == 'backward':
|
||||
tensors_fa2_pt = [(pt_qg, fa2_qg),
|
||||
(pt_kg, fa2_kg),
|
||||
(pt_vg, fa2_vg)]
|
||||
else: # 'all'
|
||||
tensors_fa2_pt = [(pt_o, fa2_o),
|
||||
(pt_qg, fa2_qg),
|
||||
(pt_kg, fa2_kg),
|
||||
(pt_vg, fa2_vg)]
|
||||
|
||||
for pt, fa2 in tensors_fa2_pt:
|
||||
diff = pt - fa2
|
||||
abs_diff = torch.abs(diff)
|
||||
results['FA2 vs PT']['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results['FA2 vs PT']['sum_abs'] += torch.sum(torch.abs(pt)).item()
|
||||
results['FA2 vs PT']['max_diff'] = max(results['FA2 vs PT']['max_diff'], torch.max(abs_diff).item())
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Calculate total elements based on test mode and error mode
|
||||
if test_mode == 'forward_only':
|
||||
total_elements = b * h * n * d * num_iterations
|
||||
else: # 'forward_backward'
|
||||
total_elements = b * h * n * d * num_iterations * (1 if error_mode == 'output' else 3 if error_mode == 'backward' else 4)
|
||||
|
||||
for name, data in results.items():
|
||||
avg_diff = data['sum_diff'] / total_elements
|
||||
max_diff = data['max_diff']
|
||||
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
|
||||
|
||||
return results
|
||||
|
||||
def generate_error_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward'):
|
||||
seq_lengths = [768 * (2**i) for i in range(1)]
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print(f"ATTENTION ERROR COMPARISON TABLE (b={b}, h={h}, d={d}, mean={mean}, std={std})")
|
||||
print(f"Mode: {error_mode}, Test: {test_mode}")
|
||||
print(f"{'='*80}")
|
||||
|
||||
# Print header
|
||||
print(f"{'Seq Length':<12} | {'FA2 vs PT Avg':<15} | {'FA2 vs PT Max':<15}")
|
||||
print(f"{'-'*12} | {'-'*15} | {'-'*15}")
|
||||
|
||||
for n in seq_lengths:
|
||||
results = check_correctness(b, h, n, d, mean, std, error_mode=error_mode, test_mode=test_mode)
|
||||
|
||||
fa2_pt_avg = results['FA2 vs PT']['avg_diff']
|
||||
fa2_pt_max = results['FA2 vs PT']['max_diff']
|
||||
|
||||
# Print row
|
||||
print(f"{n:<12} | {fa2_pt_avg:<15.6e} | {fa2_pt_max:<15.6e}")
|
||||
|
||||
print(f"{'='*80}\n")
|
||||
|
||||
# fix random seed
|
||||
torch.manual_seed(0)
|
||||
|
||||
# Example usage
|
||||
b, h, d = 2, 2, 64
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
# Test forward only
|
||||
generate_error_tables(b, h, d, mean, std, error_mode='output', test_mode='forward_only')
|
||||
|
||||
# Test forward and backward
|
||||
generate_error_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward')
|
||||
|
||||
print("Attention error comparison completed.")
|
||||
@@ -0,0 +1,175 @@
|
||||
import torch
|
||||
from flash_attn_interface import flash_attn_func
|
||||
from st_attn import mha_forward, mha_backward
|
||||
import random
|
||||
from tqdm import tqdm
|
||||
import matplotlib.pyplot as plt
|
||||
import numpy as np
|
||||
|
||||
def pytorch_test(Q, K, V, dO):
|
||||
q_ = Q.to(torch.float64).requires_grad_()
|
||||
k_ = K.to(torch.float64).requires_grad_()
|
||||
v_ = V.to(torch.float64).requires_grad_()
|
||||
dO_ = dO.to(torch.float64)
|
||||
|
||||
# manual pytorch implementation of scaled dot product attention
|
||||
QK = torch.matmul(q_, k_.transpose(-2, -1))
|
||||
QK /= (q_.size(-1) ** 0.5)
|
||||
|
||||
# Causal mask removed since causal is always false
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v_)
|
||||
|
||||
output.backward(dO_)
|
||||
|
||||
q_grad = q_.grad
|
||||
k_grad = k_.grad
|
||||
v_grad = v_.grad
|
||||
|
||||
return output, q_grad, k_grad, v_grad
|
||||
|
||||
def fa2_test(Q, K, V, dO):
|
||||
Q.requires_grad = True
|
||||
K.requires_grad = True
|
||||
V.requires_grad = True
|
||||
output = torch.nn.functional.scaled_dot_product_attention(Q, K, V, is_causal=False)
|
||||
output.backward(dO)
|
||||
|
||||
return output, Q.grad, K.grad, V.grad
|
||||
|
||||
|
||||
def mha_kernel_test(Q, K, V, dO, mode):
|
||||
Q.requires_grad = True
|
||||
K.requires_grad = True
|
||||
V.requires_grad = True
|
||||
|
||||
o, l_vec = mha_forward(Q, K, V)
|
||||
|
||||
if mode == 'forward_only':
|
||||
return o, None, None, None
|
||||
else: # 'forward_backward'
|
||||
qg, kg, vg = mha_backward(Q, K, V, o, l_vec, dO)
|
||||
return o, qg, kg, vg
|
||||
|
||||
def generate_tensor(shape, mean, std, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
|
||||
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
|
||||
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
|
||||
|
||||
return scaled_tensor.contiguous()
|
||||
|
||||
def check_correctness(b, h, n, d, mean, std, num_iterations=100, error_mode='all', test_mode='forward_backward'):
|
||||
results = {
|
||||
'MHA vs PT': {'sum_diff': 0, 'sum_abs': 0, 'max_diff': 0},
|
||||
'FA2 vs PT': {'sum_diff': 0, 'sum_abs': 0, 'max_diff': 0},
|
||||
}
|
||||
|
||||
for _ in range(num_iterations):
|
||||
torch.manual_seed(0)
|
||||
|
||||
Q = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
K = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
V = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
dO = generate_tensor((b, h, n, d), mean, std, torch.bfloat16, 'cuda')
|
||||
|
||||
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, dO)
|
||||
fa2_o, fa2_qg, fa2_kg, fa2_vg = fa2_test(Q, K, V, dO)
|
||||
|
||||
if test_mode == 'forward_only':
|
||||
mha_o, _, _, _ = mha_kernel_test(Q, K, V, dO, 'forward_only')
|
||||
tensors_mha_pt = [(pt_o, mha_o)]
|
||||
tensors_fa2_pt = [(pt_o, fa2_o)]
|
||||
else: # 'forward_backward'
|
||||
mha_o, mha_qg, mha_kg, mha_vg = mha_kernel_test(Q, K, V, dO, 'forward_backward')
|
||||
|
||||
if error_mode == 'output':
|
||||
tensors_mha_pt = [(pt_o, mha_o)]
|
||||
tensors_fa2_pt = [(pt_o, fa2_o)]
|
||||
elif error_mode == 'backward':
|
||||
tensors_mha_pt = [(pt_qg, mha_qg),
|
||||
(pt_kg, mha_kg),
|
||||
(pt_vg, mha_vg)]
|
||||
tensors_fa2_pt = [(pt_qg, fa2_qg),
|
||||
(pt_kg, fa2_kg),
|
||||
(pt_vg, fa2_vg)]
|
||||
else: # 'all'
|
||||
tensors_mha_pt = [(pt_o, mha_o),
|
||||
(pt_qg, mha_qg),
|
||||
(pt_kg, mha_kg),
|
||||
(pt_vg, mha_vg)]
|
||||
tensors_fa2_pt = [(pt_o, fa2_o),
|
||||
(pt_qg, fa2_qg),
|
||||
(pt_kg, fa2_kg),
|
||||
(pt_vg, fa2_vg)]
|
||||
|
||||
for pt, mha in tensors_mha_pt:
|
||||
diff = pt - mha
|
||||
abs_diff = torch.abs(diff)
|
||||
results['MHA vs PT']['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results['MHA vs PT']['sum_abs'] += torch.sum(torch.abs(pt)).item()
|
||||
results['MHA vs PT']['max_diff'] = max(results['MHA vs PT']['max_diff'], torch.max(abs_diff).item())
|
||||
|
||||
for pt, fa2 in tensors_fa2_pt:
|
||||
diff = pt - fa2
|
||||
abs_diff = torch.abs(diff)
|
||||
results['FA2 vs PT']['sum_diff'] += torch.sum(abs_diff).item()
|
||||
results['FA2 vs PT']['sum_abs'] += torch.sum(torch.abs(pt)).item()
|
||||
results['FA2 vs PT']['max_diff'] = max(results['FA2 vs PT']['max_diff'], torch.max(abs_diff).item())
|
||||
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
# Calculate total elements based on test mode and error mode
|
||||
if test_mode == 'forward_only':
|
||||
total_elements = b * h * n * d * num_iterations
|
||||
else: # 'forward_backward'
|
||||
total_elements = b * h * n * d * num_iterations * (1 if error_mode == 'output' else 3 if error_mode == 'backward' else 4)
|
||||
|
||||
for name, data in results.items():
|
||||
avg_diff = data['sum_diff'] / total_elements
|
||||
max_diff = data['max_diff']
|
||||
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
|
||||
|
||||
return results
|
||||
|
||||
def generate_error_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward'):
|
||||
seq_lengths = [768 * (2**i) for i in range(1)]
|
||||
|
||||
print(f"\n{'='*80}")
|
||||
print(f"MHA ERROR COMPARISON TABLE (b={b}, h={h}, d={d}, mean={mean}, std={std})")
|
||||
print(f"Mode: {error_mode}, Test: {test_mode}")
|
||||
print(f"{'='*80}")
|
||||
|
||||
# Print header
|
||||
print(f"{'Seq Length':<12} | {'MHA vs PT Avg':<15} | {'MHA vs PT Max':<15} | {'FA2 vs PT Avg':<15} | {'FA2 vs PT Max':<15}")
|
||||
print(f"{'-'*12} | {'-'*15} | {'-'*15} | {'-'*15} | {'-'*15}")
|
||||
|
||||
for n in seq_lengths:
|
||||
results = check_correctness(b, h, n, d, mean, std, error_mode=error_mode, test_mode=test_mode)
|
||||
|
||||
mha_pt_avg = results['MHA vs PT']['avg_diff']
|
||||
mha_pt_max = results['MHA vs PT']['max_diff']
|
||||
fa2_pt_avg = results['FA2 vs PT']['avg_diff']
|
||||
fa2_pt_max = results['FA2 vs PT']['max_diff']
|
||||
|
||||
# Print row
|
||||
print(f"{n:<12} | {mha_pt_avg:<15.6e} | {mha_pt_max:<15.6e} | {fa2_pt_avg:<15.6e} | {fa2_pt_max:<15.6e}")
|
||||
|
||||
print(f"{'='*80}\n")
|
||||
|
||||
# fix random seed
|
||||
torch.manual_seed(0)
|
||||
|
||||
# Example usage
|
||||
b, h, d = 2, 2, 64
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
# Test forward only
|
||||
generate_error_tables(b, h, d, mean, std, error_mode='output', test_mode='forward_only')
|
||||
|
||||
# Test forward and backward
|
||||
generate_error_tables(b, h, d, mean, std, error_mode='all', test_mode='forward_backward')
|
||||
|
||||
print("MHA attention error comparison completed.")
|
||||
@@ -2,27 +2,28 @@ import torch
|
||||
from flex_sta_ref import get_sliding_tile_attention_mask
|
||||
from st_attn import sliding_tile_attention
|
||||
from torch.nn.attention.flex_attention import flex_attention
|
||||
# from flash_attn_interface import flash_attn_func
|
||||
from tqdm import tqdm
|
||||
|
||||
flex_attention = torch.compile(flex_attention, dynamic=False)
|
||||
|
||||
|
||||
def flex_test(Q, K, V, kernel_size):
|
||||
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (36, 48, 48), 39, 'cuda', 0)
|
||||
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (18, 48, 80), 0, 'cuda', 0)
|
||||
output = flex_attention(Q, K, V, block_mask=mask)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def h100_fwd_kernel_test(Q, K, V, kernel_size):
|
||||
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 39, False)
|
||||
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
|
||||
return o
|
||||
|
||||
|
||||
def generate_tensor(shape, mean, std, dtype, device):
|
||||
tensor = torch.randn(shape, dtype=dtype, device=device)
|
||||
|
||||
magnitude = torch.linalg.norm(tensor, dim=-1, keepdim=True)
|
||||
magnitude = torch.norm(tensor, dim=-1, keepdim=True)
|
||||
scaled_tensor = tensor * (torch.randn(magnitude.shape, dtype=dtype, device=device) * std + mean) / magnitude
|
||||
|
||||
return scaled_tensor.contiguous()
|
||||
@@ -36,7 +37,7 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
|
||||
'max_diff': 0
|
||||
},
|
||||
}
|
||||
kernel_size_ls = [(6, 1, 6), (6, 6, 1)]
|
||||
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
|
||||
from tqdm import tqdm
|
||||
for kernel_size in tqdm(kernel_size_ls):
|
||||
for _ in range(num_iterations):
|
||||
@@ -71,25 +72,14 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
|
||||
return results
|
||||
|
||||
|
||||
def generate_error_graphs(b, h, d, causal, mean, std, error_mode='all'):
|
||||
seq_lengths = [82944]
|
||||
|
||||
tk_avg_errors, tk_max_errors = [], []
|
||||
|
||||
for n in tqdm(seq_lengths, desc="Generating error data"):
|
||||
results = check_correctness(b, h, n, d, causal, mean, std, error_mode=error_mode)
|
||||
|
||||
tk_avg_errors.append(results['TK vs FLEX']['avg_diff'])
|
||||
tk_max_errors.append(results['TK vs FLEX']['max_diff'])
|
||||
|
||||
|
||||
# Example usage
|
||||
b, h, d = 2, 24, 128
|
||||
n = 69120 # Sequence length
|
||||
causal = False
|
||||
mean = 1e-1
|
||||
std = 10
|
||||
|
||||
for mode in ['output']:
|
||||
generate_error_graphs(b, h, d, causal, mean, std, error_mode=mode)
|
||||
|
||||
print("Error graphs generated and saved for all modes.")
|
||||
# Run correctness check directly
|
||||
results = check_correctness(b, h, n, d, causal, mean, std, error_mode='output')
|
||||
print(f"Average difference: {results['TK vs FLEX']['avg_diff']}")
|
||||
print(f"Maximum difference: {results['TK vs FLEX']['max_diff']}")
|
||||
@@ -0,0 +1,27 @@
|
||||
#include <torch/extension.h>
|
||||
#include <ATen/ATen.h>
|
||||
|
||||
#include <vector>
|
||||
#include <cuda_fp16.h>
|
||||
#include <cuda_bf16.h>
|
||||
|
||||
#include <cuda_runtime.h>
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE
|
||||
extern std::vector<torch::Tensor> block_sparse_attention_forward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor q2k_block_sparse_index, torch::Tensor q2k_block_sparse_num
|
||||
);
|
||||
extern std::vector<torch::Tensor> block_sparse_attention_backward(
|
||||
torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o, torch::Tensor l_vec, torch::Tensor og, torch::Tensor k2q_block_sparse_index, torch::Tensor k2q_block_sparse_num
|
||||
);
|
||||
#endif
|
||||
|
||||
|
||||
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
|
||||
m.doc() = "Video Sparse Attention Kernels"; // optional module docstring
|
||||
|
||||
#ifdef TK_COMPILE_BLOCK_SPARSE
|
||||
m.def("block_sparse_fwd", torch::wrap_pybind_function(block_sparse_attention_forward), "block sparse attention");
|
||||
m.def("block_sparse_bwd", torch::wrap_pybind_function(block_sparse_attention_backward), "block sparse attention backward");
|
||||
#endif
|
||||
}
|
||||
@@ -0,0 +1,470 @@
|
||||
import math
|
||||
|
||||
import torch
|
||||
from torch.utils.checkpoint import detach_variable
|
||||
from typing import Tuple
|
||||
try:
|
||||
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
|
||||
except ImportError:
|
||||
block_sparse_fwd = None
|
||||
block_sparse_bwd = None
|
||||
|
||||
|
||||
BLOCK_M = 64
|
||||
BLOCK_N = 64
|
||||
|
||||
def video_sparse_attn(q, k, v, topk, block_size, compress_attn_weight=None):
|
||||
"""
|
||||
q: [batch_size, num_heads, seq_len, head_dim]
|
||||
k: [batch_size, num_heads, seq_len, head_dim]
|
||||
v: [batch_size, num_heads, seq_len, head_dim]
|
||||
topk: int
|
||||
block_size: int or tuple of 3 ints
|
||||
video_shape: tuple of (T, H, W)
|
||||
compress_attn_weight: [batch_size, num_heads, seq_len, head_dim]
|
||||
select_attn_weight: [batch_size, num_heads, seq_len, head_dim]
|
||||
|
||||
V1 of sparse attention. Include compress attn and sparse attn branch, use average pooling to compress.
|
||||
Assume q, k, v is flattened in this way: [batch_size, num_heads, T//block_size[0], H//block_size[1], W//block_size[2], block_size[0], block_size[1], block_size[2]]
|
||||
"""
|
||||
|
||||
if isinstance(block_size, int):
|
||||
block_size = (block_size, block_size, block_size)
|
||||
|
||||
block_elements = block_size[0] * block_size[1] * block_size[2]
|
||||
assert block_elements % 64 == 0 and block_elements >= 64
|
||||
assert q.shape[2] % block_elements == 0
|
||||
batch_size, num_heads, seq_len, head_dim = q.shape
|
||||
# compress attn
|
||||
q_compress = q.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).mean(dim=3)
|
||||
k_compress = k.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).mean(dim=3)
|
||||
v_compress = v.view(batch_size, num_heads, seq_len // block_elements,
|
||||
block_elements, head_dim).mean(dim=3)
|
||||
|
||||
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
|
||||
v_compress)
|
||||
|
||||
output_compress = output_compress.view(batch_size, num_heads,
|
||||
seq_len // block_elements, 1,
|
||||
head_dim)
|
||||
output_compress = output_compress.repeat(1, 1, 1, block_elements,
|
||||
1).view(batch_size, num_heads,
|
||||
seq_len, head_dim)
|
||||
|
||||
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num = generate_topk_block_sparse_pattern(
|
||||
block_attn_score, topk)
|
||||
|
||||
output_select = block_sparse_attn(q, k, v, q2k_block_sparse_index,
|
||||
q2k_block_sparse_num,
|
||||
k2q_block_sparse_index,
|
||||
k2q_block_sparse_num)
|
||||
|
||||
if compress_attn_weight is not None:
|
||||
final_output = output_compress * compress_attn_weight + output_select
|
||||
else:
|
||||
final_output = output_compress + output_select
|
||||
return final_output
|
||||
|
||||
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
QK = torch.matmul(q, k.transpose(-2, -1))
|
||||
QK /= (q.size(-1)**0.5)
|
||||
|
||||
# Causal mask removed since causal is always false
|
||||
|
||||
QK = torch.nn.functional.softmax(QK, dim=-1)
|
||||
output = torch.matmul(QK, v)
|
||||
return output, QK
|
||||
|
||||
def generate_topk_block_sparse_pattern(block_attn_score: torch.Tensor,
|
||||
topk: int):
|
||||
"""
|
||||
Generate a block sparse pattern where each q block attends to exactly topk kv blocks,
|
||||
based on the provided attention scores.
|
||||
|
||||
Args:
|
||||
block_attn_score: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
Attention scores between query and key blocks
|
||||
topk: int
|
||||
Number of kv blocks each q block attends to
|
||||
|
||||
Returns:
|
||||
q2k_block_sparse_index: [bs, h, num_q_blocks, topk]
|
||||
Contains the indices of kv blocks that each q block attends to.
|
||||
q2k_block_sparse_num: [bs, h, num_q_blocks]
|
||||
Contains the number of kv blocks that each q block attends to (all equal to topk).
|
||||
k2q_block_sparse_index: [bs, h, num_kv_blocks, max_q_per_kv]
|
||||
Contains the indices of q blocks that attend to each kv block.
|
||||
k2q_block_sparse_num: [bs, h, num_kv_blocks]
|
||||
Contains the number of q blocks that attend to each kv block.
|
||||
"""
|
||||
device = block_attn_score.device
|
||||
# Extract dimensions from block_attn_score
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_attn_score.shape
|
||||
|
||||
sorted_result = torch.sort(block_attn_score, dim=-1, descending=True)
|
||||
|
||||
sorted_indice = sorted_result.indices
|
||||
|
||||
q2k_block_sparse_index, _ = torch.sort(sorted_indice[:, :, :, :topk],
|
||||
dim=-1)
|
||||
q2k_block_sparse_index = q2k_block_sparse_index.to(dtype=torch.int32)
|
||||
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks),
|
||||
topk,
|
||||
device=device,
|
||||
dtype=torch.int32)
|
||||
|
||||
block_map = topk_index_to_map(q2k_block_sparse_index,
|
||||
num_kv_blocks,
|
||||
transpose_map=True)
|
||||
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(
|
||||
block_map.transpose(2, 3))
|
||||
|
||||
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
|
||||
@torch._dynamo.disable
|
||||
def block_sparse_attn(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
"""
|
||||
Differentiable block sparse attention function.
|
||||
|
||||
Args:
|
||||
q: Query tensor [batch_size, num_heads, seq_len_q, head_dim]
|
||||
k: Key tensor [batch_size, num_heads, seq_len_kv, head_dim]
|
||||
v: Value tensor [batch_size, num_heads, seq_len_kv, head_dim]
|
||||
q2k_block_sparse_index: Indices for query-to-key sparse blocks
|
||||
q2k_block_sparse_num: Number of sparse blocks for each query block
|
||||
k2q_block_sparse_index: Indices for key-to-query sparse blocks (for backward pass)
|
||||
k2q_block_sparse_num: Number of sparse blocks for each key block (for backward pass)
|
||||
|
||||
Returns:
|
||||
output: Attention output tensor [batch_size, num_heads, seq_len_q, head_dim]
|
||||
"""
|
||||
return BlockSparseAttentionFunction.apply(
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
)
|
||||
|
||||
def block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num):
|
||||
"""
|
||||
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks].
|
||||
[*, *, i, j] = 1 means the i-th q block should attend to the j-th kv block.
|
||||
"""
|
||||
# assert all elements in q2k_block_sparse_num can be devisible by 2
|
||||
o, lse = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
return o, lse
|
||||
|
||||
def block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
grad_output = grad_output.contiguous()
|
||||
grad_q, grad_k, grad_v = block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
return grad_q, grad_k, grad_v
|
||||
|
||||
## pytorch sdpa version of block sparse ##
|
||||
import triton
|
||||
import triton.language as tl
|
||||
|
||||
@triton.jit
|
||||
def index_to_mask_kernel(
|
||||
q2k_block_sparse_index_ptr,
|
||||
q2k_block_sparse_num_ptr,
|
||||
mask_ptr,
|
||||
batch_size: tl.constexpr,
|
||||
num_heads: tl.constexpr,
|
||||
num_q_blocks: tl.constexpr,
|
||||
num_k_blocks: tl.constexpr,
|
||||
max_kv_blocks: tl.constexpr,
|
||||
BLOCK_Q: tl.constexpr,
|
||||
BLOCK_K: tl.constexpr,
|
||||
):
|
||||
bh, q, id = tl.program_id(0).to(tl.int64), tl.program_id(1).to(tl.int64), tl.program_id(2).to(tl.int64)
|
||||
b = bh // num_heads
|
||||
h = bh % num_heads
|
||||
|
||||
num_valid_blocks = tl.load(q2k_block_sparse_num_ptr + b * num_heads * num_q_blocks + h * num_q_blocks + q)
|
||||
|
||||
if num_valid_blocks <= id:
|
||||
return
|
||||
k = tl.load(q2k_block_sparse_index_ptr + b * num_heads * num_q_blocks * max_kv_blocks + h * num_q_blocks * max_kv_blocks + q * max_kv_blocks + id)
|
||||
|
||||
full_mask = (tl.arange(0, BLOCK_Q)[:, None] < BLOCK_Q) & (tl.arange(0, BLOCK_K)[None, :] < BLOCK_K)
|
||||
|
||||
q_lengths = num_q_blocks * BLOCK_Q
|
||||
k_lengths = num_k_blocks * BLOCK_K
|
||||
mask_ptr_base = mask_ptr + b * num_heads * q_lengths * k_lengths + h * q_lengths * k_lengths + q * BLOCK_Q * k_lengths + k * BLOCK_K
|
||||
|
||||
tl.store(mask_ptr_base + tl.arange(0, BLOCK_Q)[:, None] * k_lengths + tl.arange(0, BLOCK_K)[None, :], full_mask)
|
||||
|
||||
def index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, BLOCK_Q, BLOCK_K, num_k_blocks):
|
||||
"""
|
||||
Convert block sparse indices to a mask.
|
||||
|
||||
Args:
|
||||
q2k_block_sparse_index: Indices for query-to-key sparse blocks
|
||||
q2k_block_sparse_num: Number of sparse blocks for each query block
|
||||
|
||||
Returns:
|
||||
mask: Block sparse mask tensor
|
||||
"""
|
||||
batch_size, num_heads, num_q_blocks, max_kv_blocks = q2k_block_sparse_index.shape
|
||||
assert q2k_block_sparse_num.shape == (batch_size, num_heads, num_q_blocks)
|
||||
|
||||
mask = torch.zeros((batch_size, num_heads, num_q_blocks * BLOCK_Q, num_k_blocks * BLOCK_K), dtype=torch.bool, device=q2k_block_sparse_index.device)
|
||||
|
||||
grid = (batch_size * num_heads, num_q_blocks, max_kv_blocks)
|
||||
index_to_mask_kernel[grid](
|
||||
q2k_block_sparse_index,
|
||||
q2k_block_sparse_num,
|
||||
mask,
|
||||
batch_size,
|
||||
num_heads,
|
||||
num_q_blocks,
|
||||
num_k_blocks,
|
||||
max_kv_blocks,
|
||||
BLOCK_Q=BLOCK_Q,
|
||||
BLOCK_K=BLOCK_K,
|
||||
)
|
||||
|
||||
return mask
|
||||
|
||||
@triton.jit
|
||||
def topk_index_to_map_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
topk: tl.constexpr,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
for i in tl.static_range(topk):
|
||||
index = tl.load(index_ptr_base + i * index_kv_stride)
|
||||
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
|
||||
|
||||
@triton.jit
|
||||
def map_to_index_kernel(
|
||||
map_ptr,
|
||||
index_ptr,
|
||||
index_num_ptr,
|
||||
map_bs_stride,
|
||||
map_h_stride,
|
||||
map_q_stride,
|
||||
map_kv_stride,
|
||||
index_bs_stride,
|
||||
index_h_stride,
|
||||
index_q_stride,
|
||||
index_kv_stride,
|
||||
index_num_bs_stride,
|
||||
index_num_h_stride,
|
||||
index_num_q_stride,
|
||||
num_kv_blocks: tl.constexpr,
|
||||
):
|
||||
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
|
||||
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
|
||||
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
|
||||
|
||||
num = 0
|
||||
for i in tl.static_range(num_kv_blocks):
|
||||
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
|
||||
if map_entry:
|
||||
tl.store(index_ptr_base + num * index_kv_stride, i)
|
||||
num += 1
|
||||
|
||||
tl.store(
|
||||
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
|
||||
q * index_num_q_stride, num)
|
||||
|
||||
def topk_index_to_map(index: torch.Tensor,
|
||||
num_kv_blocks: int,
|
||||
transpose_map: bool = False):
|
||||
"""
|
||||
Convert topk indices to a map.
|
||||
|
||||
Args:
|
||||
index: [bs, h, num_q_blocks, topk]
|
||||
The topk indices tensor.
|
||||
num_kv_blocks: int
|
||||
The number of key-value blocks in the block_map returned
|
||||
transpose_map: bool
|
||||
If True, the block_map will be transposed on the final two dimensions.
|
||||
|
||||
Returns:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
A binary map where 1 indicates that the q block attends to the kv block.
|
||||
"""
|
||||
bs, h, num_q_blocks, topk = index.shape
|
||||
|
||||
if transpose_map is False:
|
||||
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
else:
|
||||
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
|
||||
dtype=torch.bool,
|
||||
device=index.device)
|
||||
block_map = block_map.transpose(2, 3)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
topk_index_to_map_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
topk=topk,
|
||||
)
|
||||
|
||||
return block_map
|
||||
|
||||
def map_to_index(block_map: torch.Tensor):
|
||||
"""
|
||||
Convert a block map to indices and counts.
|
||||
|
||||
Args:
|
||||
block_map: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The block map tensor.
|
||||
|
||||
Returns:
|
||||
index: [bs, h, num_q_blocks, num_kv_blocks]
|
||||
The indices of the blocks.
|
||||
index_num: [bs, h, num_q_blocks]
|
||||
The number of blocks for each q block.
|
||||
"""
|
||||
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
|
||||
|
||||
index = torch.full((block_map.shape),
|
||||
-1,
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
index_num = torch.empty((bs, h, num_q_blocks),
|
||||
dtype=torch.int32,
|
||||
device=block_map.device)
|
||||
|
||||
grid = (bs, h, num_q_blocks)
|
||||
map_to_index_kernel[grid](
|
||||
block_map,
|
||||
index,
|
||||
index_num,
|
||||
block_map.stride(0),
|
||||
block_map.stride(1),
|
||||
block_map.stride(2),
|
||||
block_map.stride(3),
|
||||
index.stride(0),
|
||||
index.stride(1),
|
||||
index.stride(2),
|
||||
index.stride(3),
|
||||
index_num.stride(0),
|
||||
index_num.stride(1),
|
||||
index_num.stride(2),
|
||||
num_kv_blocks=num_kv_blocks,
|
||||
)
|
||||
|
||||
return index, index_num
|
||||
|
||||
class BlockSparseAttentionFunction(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
|
||||
o, lse = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
|
||||
ctx.save_for_backward(q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num)
|
||||
return o
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num = ctx.saved_tensors
|
||||
grad_q, grad_k, grad_v = block_sparse_attention_backward(
|
||||
q, k, v, o, lse, grad_output, k2q_block_sparse_index, k2q_block_sparse_num
|
||||
)
|
||||
return grad_q, grad_k, grad_v, None, None, None, None
|
||||
|
||||
|
||||
class DummyOperator(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, x):
|
||||
return x
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
return grad_output
|
||||
|
||||
class CheckpointSDPA(torch.autograd.Function):
|
||||
@staticmethod
|
||||
def forward(ctx, obj, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k):
|
||||
"""Forward pass."""
|
||||
with torch.no_grad():
|
||||
mask = index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k, k.shape[2] // block_k)
|
||||
outputs = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
|
||||
ctx.save_for_backward(*detach_variable((q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)))
|
||||
ctx.block_q = block_q
|
||||
ctx.block_k = block_k
|
||||
# the obj is passed in, then it can access the saved input
|
||||
# tensors later for recomputation
|
||||
obj.ctx = ctx
|
||||
return outputs
|
||||
|
||||
@staticmethod
|
||||
def backward(ctx, grad_output):
|
||||
"""Backward pass."""
|
||||
inputs = ctx.saved_tensors
|
||||
output = ctx.output
|
||||
torch.autograd.backward(output, grad_output)
|
||||
ctx.output = None
|
||||
grads = tuple(inp.grad for inp in inputs)
|
||||
return (None, ) + grads + (None, None)
|
||||
|
||||
|
||||
class BlockSparseAttnTorch:
|
||||
def __init__(self):
|
||||
self.ctx = None
|
||||
|
||||
def recompute_mask(self, _):
|
||||
recomputed_mask = index_to_mask(self.q2k_block_sparse_index, self.q2k_block_sparse_num, self.block_q, self.block_k, self.num_kv_blocks)
|
||||
mask_size = recomputed_mask.untyped_storage().size()
|
||||
self.mask.untyped_storage().resize_(mask_size)
|
||||
self.mask.untyped_storage().copy_(recomputed_mask.untyped_storage())
|
||||
|
||||
def recompute(self, _):
|
||||
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num = self.ctx.saved_tensors
|
||||
block_q = self.ctx.block_q
|
||||
block_k = self.ctx.block_k
|
||||
mask = index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k, k.shape[2] // block_k)
|
||||
with torch.enable_grad():
|
||||
output = torch.nn.functional.scaled_dot_product_attention(q, k, v, attn_mask=mask)
|
||||
self.ctx.output = output
|
||||
self.ctx = None
|
||||
|
||||
@torch._dynamo.disable
|
||||
def forward(self, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k):
|
||||
"""
|
||||
Differentiable block sparse attention function using PyTorch.
|
||||
|
||||
Args:
|
||||
q: Query tensor [batch_size, num_heads, seq_len_q, head_dim]
|
||||
k: Key tensor [batch_size, num_heads, seq_len_kv, head_dim]
|
||||
v: Value tensor [batch_size, num_heads, seq_len_kv, head_dim]
|
||||
q2k_block_sparse_index: Indices for query-to-key sparse blocks
|
||||
q2k_block_sparse_num: Number of sparse blocks for each query block
|
||||
block_q: Block size for query
|
||||
block_k: Block size for key-value
|
||||
|
||||
Returns:
|
||||
output: Attention output tensor [batch_size, num_heads, seq_len_q, head_dim]
|
||||
"""
|
||||
|
||||
output = CheckpointSDPA.apply(
|
||||
self, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, block_q, block_k
|
||||
)
|
||||
|
||||
o = DummyOperator.apply(output)
|
||||
o.register_hook(self.recompute)
|
||||
return o
|
||||
File diff suppressed because it is too large
Load Diff
@@ -288,7 +288,7 @@ Sequence parallelism splits sequences across devices:
|
||||
|
||||
```python
|
||||
# Distributed attention for long sequences
|
||||
from fastvideo.v1.attention import DistributedAttention
|
||||
from fastvideo.v1.layers.attention import DistributedAttention
|
||||
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
|
||||
@@ -96,8 +96,8 @@ Replace standard attention with FastVideo's optimized attention:
|
||||
|
||||
```python
|
||||
# Local attention patterns
|
||||
from fastvideo.v1.attention import LocalAttention
|
||||
from fastvideo.v1.attention.backends.abstract import _Backend
|
||||
from fastvideo.v1.layers.attention import LocalAttention
|
||||
from fastvideo.v1.layers.attention.backends.abstract import _Backend
|
||||
self.attn = LocalAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
@@ -108,7 +108,7 @@ self.attn = LocalAttention(
|
||||
)
|
||||
|
||||
# Distributed attention for long sequences
|
||||
from fastvideo.v1.attention import DistributedAttention
|
||||
from fastvideo.v1.layers.attention import DistributedAttention
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=head_dim,
|
||||
|
||||
@@ -57,8 +57,9 @@ Run the script with:
|
||||
python example.py
|
||||
```
|
||||
|
||||
The generated video will be saved in the current directory under `my_videos/`.
|
||||
The generated video will be saved in the current directory under `my_videos/`
|
||||
|
||||
More inference example scripts can be found in `scripts/inference/`
|
||||
## Available Models
|
||||
|
||||
Please see the [support matrix](#support-matrix) for the list of supported models and their available optimizations.
|
||||
@@ -79,7 +80,6 @@ def main():
|
||||
sampling_param = SamplingParam.from_pretrained(model_name)
|
||||
sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
|
||||
sampling_param.num_frames = 107
|
||||
sampling_param.image_strength = 0.8 # How much to preserve the original image (0-1)
|
||||
|
||||
# Generate video based on the image
|
||||
prompt = "A photograph coming to life with gentle movement"
|
||||
|
||||
@@ -72,12 +72,6 @@ FastVideo will automatically detect and use `FA3` if it is installed when using
|
||||
pip install st_attn==0.0.4
|
||||
```
|
||||
|
||||
Then download STA mask strategy from Hugging Face
|
||||
|
||||
```bash
|
||||
python scripts/huggingface/download_hf.py --repo_id=FastVideo/STA_Mask_Strategy --local_dir=assets/ --repo_type=dataset
|
||||
```
|
||||
|
||||
Please see [this page](#sta-installation) for more installation instructions.
|
||||
|
||||
(optimizations-sage)=
|
||||
|
||||
@@ -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,10 @@ 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
|
||||
huggingface-cli download hunyuanvideo-community/HunyuanVideo --local-dir data/hunyuan
|
||||
bash scripts/inference/inference_hunyuan_STA.sh
|
||||
```
|
||||
|
||||
@@ -2,7 +2,7 @@ from fastvideo import VideoGenerator
|
||||
|
||||
# from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
|
||||
OUTPUT_PATH = "video_samples"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
@@ -11,7 +11,9 @@ def main():
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
# if num_gpus > 1, FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
num_gpus=2,
|
||||
use_fsdp_inference=True,
|
||||
use_cpu_offload=False
|
||||
)
|
||||
|
||||
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
|
||||
@@ -23,7 +25,7 @@ def main():
|
||||
"wide with interest. The playful yet serene atmosphere is complemented by soft "
|
||||
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
|
||||
)
|
||||
video = generator.generate_video(prompt)
|
||||
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
|
||||
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
|
||||
|
||||
# Generate another video with a different prompt, without reloading the
|
||||
@@ -34,7 +36,7 @@ def main():
|
||||
"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.")
|
||||
video2 = generator.generate_video(prompt2)
|
||||
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
from fastvideo.v1.configs.pipelines.base import PipelineConfig
|
||||
|
||||
def main():
|
||||
|
||||
|
||||
@@ -0,0 +1,45 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.v1.configs.sample import SamplingParam
|
||||
|
||||
OUTPUT_PATH = "./lora"
|
||||
def main():
|
||||
# Initialize VideoGenerator with the Wan model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
|
||||
num_gpus=2,
|
||||
lora_path="benjamin-paine/steamboat-willie-1.3b",
|
||||
lora_nickname="steamboat"
|
||||
)
|
||||
kwargs = {
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 81,
|
||||
"guidance_scale": 5.0,
|
||||
"num_inference_steps": 32,
|
||||
}
|
||||
# Generate video with LoRA style
|
||||
prompt = "steamboat willie style, golden era animation, close-up of a short fluffy monster kneeling beside a melting red candle. the mood is one of wonder and curiosity, as the monster gazes at the flame with wide eyes and open mouth. Its pose and expression convey a sense of innocence and playfulness, as if it is exploring the world around it for the first time. The use of warm colors and dramatic lighting further enhances the cozy atmosphere of the image."
|
||||
negative_prompt = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
# sampling_param=sampling_param,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
negative_prompt=negative_prompt,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
generator.set_lora_adapter(lora_nickname="flat_color", lora_path="motimalu/wan-flat-color-1.3b-v2")
|
||||
prompt = "flat color, no lineart, blending, negative space, artist:[john kafka|ponsuke kaikai|hara id 21|yoneyama mai|fuzichoco], 1girl, sakura miko, pink hair, cowboy shot, white shirt, floral print, off shoulder, outdoors, cherry blossom, tree shade, wariza, looking up, falling petals, half-closed eyes, white sky, clouds, live2d animation, upper body, high quality cinematic video of a woman sitting under a sakura tree. Dreamy and lonely, the camera close-ups on the face of the woman as she turns towards the viewer. The Camera is steady, This is a cowboy shot. The animation is smooth and fluid."
|
||||
negative_prompt = "bad quality video,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
negative_prompt=negative_prompt,
|
||||
**kwargs
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,5 @@
|
||||
# STA Mask Search Examples
|
||||
|
||||
```bash
|
||||
bash examples/inference/sta_mask_search/inference_wan_sta.sh
|
||||
```
|
||||
@@ -0,0 +1,39 @@
|
||||
#!/bin/bash
|
||||
|
||||
export FASTVIDEO_ATTENTION_CONFIG=assets/mask_strategy_wan.json
|
||||
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)
|
||||
gpu_ids=$(seq 0 $((num_gpu-1)))
|
||||
skip_time_steps=12
|
||||
|
||||
output_path="inference_results/sta/mask_search_full"
|
||||
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 \
|
||||
--output_path $output_path \
|
||||
--STA_mode $STA_mode &
|
||||
sleep 1
|
||||
done
|
||||
wait
|
||||
echo "STA searching completed"
|
||||
|
||||
output_path="inference_results/sta/mask_search_sparse"
|
||||
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 \
|
||||
--output_path $output_path \
|
||||
--STA_mode $STA_mode \
|
||||
--skip_time_steps $skip_time_steps &
|
||||
sleep 1
|
||||
done
|
||||
wait
|
||||
echo "STA tuning completed"
|
||||
|
||||
echo "All jobs completed"
|
||||
@@ -0,0 +1,63 @@
|
||||
import os
|
||||
import argparse
|
||||
from fastvideo import VideoGenerator, SamplingParam
|
||||
|
||||
def main(args):
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
# Create a video generator with a pre-trained model
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"Wan-AI/Wan2.1-T2V-14B-Diffusers",
|
||||
num_gpus=args.num_gpus, # Adjust based on your hardware
|
||||
STA_mode=args.STA_mode,
|
||||
skip_time_steps=args.skip_time_steps
|
||||
)
|
||||
|
||||
# Prompts for your video
|
||||
prompt = args.prompt
|
||||
prompt_path = args.prompt_path
|
||||
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"
|
||||
|
||||
if prompt_path is not None:
|
||||
with open(prompt_path, "r") as f:
|
||||
prompts = f.readlines()
|
||||
else:
|
||||
prompts = [prompt]
|
||||
|
||||
params = SamplingParam(
|
||||
height=args.height,
|
||||
width=args.width,
|
||||
num_frames=args.num_frames,
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
fps=args.fps,
|
||||
guidance_scale=args.guidance_scale,
|
||||
seed=args.seed,
|
||||
return_frames=True, # Also return frames from this call (defaults to False)
|
||||
output_path=args.output_path, # Controls where videos are saved
|
||||
save_video=True,
|
||||
negative_prompt=negative_prompt
|
||||
)
|
||||
|
||||
# Generate the video
|
||||
for prompt in prompts:
|
||||
video = generator.generate_video(
|
||||
prompt,
|
||||
sampling_param=params,
|
||||
)
|
||||
|
||||
if __name__ == '__main__':
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--prompt", type=str, default="A man is dancing.")
|
||||
parser.add_argument("--prompt_path", type=str, default=None)
|
||||
parser.add_argument("--height", type=int, default=768)
|
||||
parser.add_argument("--width", type=int, default=1280)
|
||||
parser.add_argument("--num_frames", type=int, default=69)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=50)
|
||||
parser.add_argument("--fps", type=int, default=16)
|
||||
parser.add_argument("--guidance_scale", type=float, default=5.0)
|
||||
parser.add_argument("--seed", type=int, default=12345)
|
||||
parser.add_argument("--output_path", type=str, default="my_videos/")
|
||||
parser.add_argument("--num_gpus", type=int, default=1)
|
||||
parser.add_argument("--STA_mode", type=str, default="STA_searching")
|
||||
parser.add_argument("--skip_time_steps", type=int, default=12)
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
@@ -68,7 +68,8 @@ def main(args):
|
||||
train_dataset = T5dataset(latents_json_path, args.vae_debug)
|
||||
text_encoder = load_text_encoder(args.model_type, args.model_path, device=device)
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
|
||||
vae.enable_tiling()
|
||||
if args.model_type != "wan":
|
||||
vae.enable_tiling()
|
||||
sampler = DistributedSampler(train_dataset, rank=local_rank, num_replicas=world_size, shuffle=True)
|
||||
train_dataloader = DataLoader(
|
||||
train_dataset,
|
||||
|
||||
@@ -33,7 +33,8 @@ def main(args):
|
||||
if not dist.is_initialized():
|
||||
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.model_path)
|
||||
vae.enable_tiling()
|
||||
if args.model_type != "wan":
|
||||
vae.enable_tiling()
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
os.makedirs(os.path.join(args.output_dir, "latent"), exist_ok=True)
|
||||
|
||||
@@ -103,13 +104,7 @@ if __name__ == "__main__":
|
||||
default=None,
|
||||
help="The output directory where the model predictions and checkpoints will be written.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
|
||||
args = parser.parse_args()
|
||||
main(args)
|
||||
|
||||
+75
-48
@@ -12,6 +12,7 @@ import torch.distributed as dist
|
||||
import wandb
|
||||
from accelerate.utils import set_seed
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from fastvideo.distill.solver import PCMFMScheduler
|
||||
from diffusers.optimization import get_scheduler
|
||||
from diffusers.utils import check_min_version
|
||||
from peft import LoraConfig
|
||||
@@ -23,7 +24,7 @@ from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
|
||||
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
|
||||
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.utils.latents_utils import normalize_dit_input
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
from fastvideo.utils.checkpoint import (resume_lora_optimizer, save_checkpoint, save_lora_checkpoint)
|
||||
from fastvideo.utils.communications import (broadcast, sp_parallel_dataloader_wrapper)
|
||||
@@ -123,13 +124,21 @@ def distill_one_step(
|
||||
noisy_model_input = sigmas * noise + (1.0 - sigmas) * model_input
|
||||
# Predict the noise residual
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": encoder_attention_mask, # B, L
|
||||
"return_dict": False,
|
||||
}
|
||||
if args.model_type == "wan":
|
||||
teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"return_dict": True,
|
||||
}
|
||||
else:
|
||||
teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": encoder_attention_mask, # B, L
|
||||
"return_dict": False,
|
||||
}
|
||||
if hunyuan_teacher_disable_cfg:
|
||||
teacher_kwargs["guidance"] = torch.tensor([1000.0],
|
||||
device=noisy_model_input.device,
|
||||
@@ -141,47 +150,70 @@ def distill_one_step(
|
||||
with torch.no_grad():
|
||||
w = distill_cfg
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
cond_teacher_output = teacher_transformer(
|
||||
noisy_model_input,
|
||||
encoder_hidden_states,
|
||||
timesteps,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
if args.model_type == "wan":
|
||||
cond_teacher_kwargs ={
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"return_dict": True,
|
||||
}
|
||||
else:
|
||||
cond_teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": encoder_attention_mask, # B, L
|
||||
"return_dict": False,
|
||||
}
|
||||
cond_teacher_output = teacher_transformer(**cond_teacher_kwargs)[0].float()
|
||||
if not_apply_cfg_solver:
|
||||
uncond_teacher_output = cond_teacher_output
|
||||
else:
|
||||
# Get teacher model prediction on noisy_latents and unconditional embedding
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
uncond_teacher_output = teacher_transformer(
|
||||
noisy_model_input,
|
||||
uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
|
||||
timesteps,
|
||||
uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
|
||||
return_dict=False,
|
||||
)[0].float()
|
||||
if args.model_type == "wan":
|
||||
uncond_teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states":uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
|
||||
"timestep": timesteps,
|
||||
"return_dict": True,
|
||||
}
|
||||
else:
|
||||
uncond_teacher_kwargs = {
|
||||
"hidden_states": noisy_model_input,
|
||||
"encoder_hidden_states":uncond_prompt_embed.unsqueeze(0).expand(bsz, -1, -1),
|
||||
"timestep": timesteps,
|
||||
"encoder_attention_mask": uncond_prompt_mask.unsqueeze(0).expand(bsz, -1),
|
||||
"return_dict": False,
|
||||
}
|
||||
|
||||
uncond_teacher_output = teacher_transformer(**uncond_teacher_kwargs)[0].float()
|
||||
|
||||
teacher_output = uncond_teacher_output + w * (cond_teacher_output - uncond_teacher_output)
|
||||
x_prev = solver.euler_step(noisy_model_input, teacher_output, index)
|
||||
|
||||
# 20.4.12. Get target LCM prediction on x_prev, w, c, t_n
|
||||
with torch.no_grad():
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
if ema_transformer is not None:
|
||||
target_pred = ema_transformer(
|
||||
x_prev.float(),
|
||||
encoder_hidden_states,
|
||||
timesteps_prev,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)[0]
|
||||
if args.model_type == "wan":
|
||||
target_pred_kwargs = {
|
||||
"hidden_states": x_prev.float(),
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep":timesteps_prev,
|
||||
"return_dict":True,
|
||||
}
|
||||
else:
|
||||
target_pred = transformer(
|
||||
x_prev.float(),
|
||||
encoder_hidden_states,
|
||||
timesteps_prev,
|
||||
encoder_attention_mask, # B, L
|
||||
return_dict=False,
|
||||
)[0]
|
||||
target_pred_kwargs = {
|
||||
"hidden_states": x_prev.float(),
|
||||
"encoder_hidden_states": encoder_hidden_states,
|
||||
"timestep":timesteps_prev,
|
||||
"encoder_attention_mask":encoder_attention_mask,
|
||||
"return_dict":False,
|
||||
}
|
||||
if ema_transformer is not None:
|
||||
target_pred = ema_transformer(**target_pred_kwargs)[0]
|
||||
else:
|
||||
target_pred = transformer(**target_pred_kwargs)[0]
|
||||
|
||||
target, end_index = solver.euler_style_multiphase_pred(x_prev, target_pred, index, multiphase, True)
|
||||
|
||||
@@ -242,7 +274,7 @@ def main(args):
|
||||
noise_random_generator = None
|
||||
|
||||
# Handle the repository creation
|
||||
if rank <= 0 and args.output_dir is not None:
|
||||
if rank == 0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weights to half-precision
|
||||
@@ -319,7 +351,9 @@ def main(args):
|
||||
teacher_transformer.requires_grad_(False)
|
||||
if args.use_ema:
|
||||
ema_transformer.requires_grad_(False)
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler(shift=args.shift)
|
||||
|
||||
noise_scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
|
||||
if args.scheduler_type == "pcm_linear_quadratic":
|
||||
linear_steps = int(noise_scheduler.config.num_train_timesteps * args.linear_range)
|
||||
sigmas = linear_quadratic_schedule(
|
||||
@@ -391,7 +425,7 @@ def main(args):
|
||||
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
@@ -493,7 +527,7 @@ def main(args):
|
||||
"phases": num_phases,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
wandb.log(
|
||||
{
|
||||
"train_loss": loss,
|
||||
@@ -637,13 +671,6 @@ if __name__ == "__main__":
|
||||
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
parser.add_argument("--num_train_epochs", type=int, default=100)
|
||||
|
||||
@@ -23,7 +23,7 @@ from tqdm.auto import tqdm
|
||||
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
|
||||
from fastvideo.distill.discriminator import Discriminator
|
||||
from fastvideo.distill.solver import EulerSolver, extract_into_tensor
|
||||
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.utils.latents_utils import normalize_dit_input
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import linear_quadratic_schedule
|
||||
from fastvideo.utils.checkpoint import (resume_lora_optimizer, resume_training_generator_discriminator, save_checkpoint,
|
||||
save_lora_checkpoint)
|
||||
@@ -296,7 +296,7 @@ def main(args):
|
||||
noise_random_generator = None
|
||||
|
||||
# Handle the repository creation
|
||||
if rank <= 0 and args.output_dir is not None:
|
||||
if rank == 0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weights to half-precision
|
||||
@@ -462,7 +462,7 @@ def main(args):
|
||||
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
@@ -559,7 +559,7 @@ def main(args):
|
||||
"step_time": f"{step_time:.2f}s",
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
wandb.log(
|
||||
{
|
||||
"generator_loss": generator_loss,
|
||||
@@ -693,13 +693,6 @@ if __name__ == "__main__":
|
||||
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
parser.add_argument("--num_train_epochs", type=int, default=100)
|
||||
|
||||
@@ -0,0 +1,486 @@
|
||||
# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import math
|
||||
from typing import Any, Dict, Optional, Tuple, Union
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
import torch.nn.functional as F
|
||||
|
||||
from diffusers.configuration_utils import ConfigMixin, register_to_config
|
||||
from diffusers.loaders import FromOriginalModelMixin, PeftAdapterMixin
|
||||
from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscale_lora_layers
|
||||
from diffusers.models.attention import FeedForward
|
||||
from diffusers.models.attention_processor import Attention
|
||||
from diffusers.models.cache_utils import CacheMixin
|
||||
from diffusers.models.embeddings import PixArtAlphaTextProjection, TimestepEmbedding, Timesteps, get_1d_rotary_pos_embed
|
||||
from diffusers.models.modeling_outputs import Transformer2DModelOutput
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
from diffusers.models.normalization import FP32LayerNorm
|
||||
|
||||
from fastvideo.models.flash_attn_no_pad import flash_attn_no_pad
|
||||
from fastvideo.utils.communications import all_gather, all_to_all_4D
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
|
||||
class WanAttnProcessor2_0:
|
||||
def __init__(self):
|
||||
if not hasattr(F, "scaled_dot_product_attention"):
|
||||
raise ImportError("WanAttnProcessor2_0 requires PyTorch 2.0. To use it, please upgrade PyTorch to 2.0.")
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
attn: Attention,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: Optional[torch.Tensor] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
rotary_emb: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
encoder_hidden_states_img = None
|
||||
if attn.add_k_proj is not None:
|
||||
# 512 is the context length of the text encoder, hardcoded for now
|
||||
image_context_length = encoder_hidden_states.shape[1] - 512
|
||||
encoder_hidden_states_img = encoder_hidden_states[:, :image_context_length]
|
||||
encoder_hidden_states = encoder_hidden_states[:, image_context_length:]
|
||||
if encoder_hidden_states is None:
|
||||
encoder_hidden_states = hidden_states
|
||||
|
||||
query = attn.to_q(hidden_states)
|
||||
key = attn.to_k(encoder_hidden_states)
|
||||
value = attn.to_v(encoder_hidden_states)
|
||||
|
||||
if attn.norm_q is not None:
|
||||
query = attn.norm_q(query)
|
||||
if attn.norm_k is not None:
|
||||
key = attn.norm_k(key)
|
||||
|
||||
query = query.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
key = key.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
value = value.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
|
||||
if rotary_emb is not None:
|
||||
|
||||
def apply_rotary_emb(hidden_states: torch.Tensor, freqs: torch.Tensor):
|
||||
x_rotated = torch.view_as_complex(hidden_states.to(torch.float64).unflatten(3, (-1, 2)))
|
||||
x_out = torch.view_as_real(x_rotated * freqs).flatten(3, 4)
|
||||
return x_out.type_as(hidden_states)
|
||||
|
||||
query = apply_rotary_emb(query, rotary_emb)
|
||||
key = apply_rotary_emb(key, rotary_emb)
|
||||
|
||||
# I2V task
|
||||
hidden_states_img = None
|
||||
if encoder_hidden_states_img is not None:
|
||||
key_img = attn.add_k_proj(encoder_hidden_states_img)
|
||||
key_img = attn.norm_added_k(key_img)
|
||||
value_img = attn.add_v_proj(encoder_hidden_states_img)
|
||||
|
||||
key_img = key_img.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
value_img = value_img.unflatten(2, (attn.heads, -1)).transpose(1, 2)
|
||||
|
||||
hidden_states_img = F.scaled_dot_product_attention(
|
||||
query, key_img, value_img, attn_mask=None, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
hidden_states_img = hidden_states_img.transpose(1, 2).flatten(2, 3)
|
||||
hidden_states_img = hidden_states_img.type_as(query)
|
||||
|
||||
hidden_states = F.scaled_dot_product_attention(
|
||||
query, key, value, attn_mask=attention_mask, dropout_p=0.0, is_causal=False
|
||||
)
|
||||
hidden_states = hidden_states.transpose(1, 2).flatten(2, 3)
|
||||
hidden_states = hidden_states.type_as(query)
|
||||
|
||||
if hidden_states_img is not None:
|
||||
hidden_states = hidden_states + hidden_states_img
|
||||
|
||||
hidden_states = attn.to_out[0](hidden_states)
|
||||
hidden_states = attn.to_out[1](hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class WanImageEmbedding(torch.nn.Module):
|
||||
def __init__(self, in_features: int, out_features: int, pos_embed_seq_len=None):
|
||||
super().__init__()
|
||||
|
||||
self.norm1 = FP32LayerNorm(in_features)
|
||||
self.ff = FeedForward(in_features, out_features, mult=1, activation_fn="gelu")
|
||||
self.norm2 = FP32LayerNorm(out_features)
|
||||
if pos_embed_seq_len is not None:
|
||||
self.pos_embed = nn.Parameter(torch.zeros(1, pos_embed_seq_len, in_features))
|
||||
else:
|
||||
self.pos_embed = None
|
||||
|
||||
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
|
||||
if self.pos_embed is not None:
|
||||
batch_size, seq_len, embed_dim = encoder_hidden_states_image.shape
|
||||
encoder_hidden_states_image = encoder_hidden_states_image.view(-1, 2 * seq_len, embed_dim)
|
||||
encoder_hidden_states_image = encoder_hidden_states_image + self.pos_embed
|
||||
|
||||
hidden_states = self.norm1(encoder_hidden_states_image)
|
||||
hidden_states = self.ff(hidden_states)
|
||||
hidden_states = self.norm2(hidden_states)
|
||||
return hidden_states
|
||||
|
||||
|
||||
class WanTimeTextImageEmbedding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
time_freq_dim: int,
|
||||
time_proj_dim: int,
|
||||
text_embed_dim: int,
|
||||
image_embed_dim: Optional[int] = None,
|
||||
pos_embed_seq_len: Optional[int] = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.timesteps_proj = Timesteps(num_channels=time_freq_dim, flip_sin_to_cos=True, downscale_freq_shift=0)
|
||||
self.time_embedder = TimestepEmbedding(in_channels=time_freq_dim, time_embed_dim=dim)
|
||||
self.act_fn = nn.SiLU()
|
||||
self.time_proj = nn.Linear(dim, time_proj_dim)
|
||||
self.text_embedder = PixArtAlphaTextProjection(text_embed_dim, dim, act_fn="gelu_tanh")
|
||||
|
||||
self.image_embedder = None
|
||||
if image_embed_dim is not None:
|
||||
self.image_embedder = WanImageEmbedding(image_embed_dim, dim, pos_embed_seq_len=pos_embed_seq_len)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
timestep: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_hidden_states_image: Optional[torch.Tensor] = None,
|
||||
):
|
||||
timestep = self.timesteps_proj(timestep)
|
||||
|
||||
time_embedder_dtype = next(iter(self.time_embedder.parameters())).dtype
|
||||
if timestep.dtype != time_embedder_dtype and time_embedder_dtype != torch.int8:
|
||||
timestep = timestep.to(time_embedder_dtype)
|
||||
temb = self.time_embedder(timestep).type_as(encoder_hidden_states)
|
||||
timestep_proj = self.time_proj(self.act_fn(temb))
|
||||
|
||||
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
|
||||
if encoder_hidden_states_image is not None:
|
||||
encoder_hidden_states_image = self.image_embedder(encoder_hidden_states_image)
|
||||
|
||||
return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image
|
||||
|
||||
|
||||
class WanRotaryPosEmbed(nn.Module):
|
||||
def __init__(
|
||||
self, attention_head_dim: int, patch_size: Tuple[int, int, int], max_seq_len: int, theta: float = 10000.0
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.attention_head_dim = attention_head_dim
|
||||
self.patch_size = patch_size
|
||||
self.max_seq_len = max_seq_len
|
||||
|
||||
h_dim = w_dim = 2 * (attention_head_dim // 6)
|
||||
t_dim = attention_head_dim - h_dim - w_dim
|
||||
|
||||
freqs = []
|
||||
for dim in [t_dim, h_dim, w_dim]:
|
||||
freq = get_1d_rotary_pos_embed(
|
||||
dim, max_seq_len, theta, use_real=False, repeat_interleave_real=False, freqs_dtype=torch.float64
|
||||
)
|
||||
freqs.append(freq)
|
||||
self.freqs = torch.cat(freqs, dim=1)
|
||||
|
||||
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.patch_size
|
||||
ppf, pph, ppw = num_frames // p_t, height // p_h, width // p_w
|
||||
|
||||
freqs = self.freqs.to(hidden_states.device)
|
||||
freqs = freqs.split_with_sizes(
|
||||
[
|
||||
self.attention_head_dim // 2 - 2 * (self.attention_head_dim // 6),
|
||||
self.attention_head_dim // 6,
|
||||
self.attention_head_dim // 6,
|
||||
],
|
||||
dim=1,
|
||||
)
|
||||
|
||||
freqs_f = freqs[0][:ppf].view(ppf, 1, 1, -1).expand(ppf, pph, ppw, -1)
|
||||
freqs_h = freqs[1][:pph].view(1, pph, 1, -1).expand(ppf, pph, ppw, -1)
|
||||
freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1)
|
||||
freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1).reshape(1, 1, ppf * pph * ppw, -1)
|
||||
return freqs
|
||||
|
||||
|
||||
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,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
# 1. Self-attention
|
||||
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
self.attn1 = Attention(
|
||||
query_dim=dim,
|
||||
heads=num_heads,
|
||||
kv_heads=num_heads,
|
||||
dim_head=dim // num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps,
|
||||
bias=True,
|
||||
cross_attention_dim=None,
|
||||
out_bias=True,
|
||||
processor=WanAttnProcessor2_0(),
|
||||
)
|
||||
|
||||
# 2. Cross-attention
|
||||
self.attn2 = Attention(
|
||||
query_dim=dim,
|
||||
heads=num_heads,
|
||||
kv_heads=num_heads,
|
||||
dim_head=dim // num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps,
|
||||
bias=True,
|
||||
cross_attention_dim=None,
|
||||
out_bias=True,
|
||||
added_kv_proj_dim=added_kv_proj_dim,
|
||||
added_proj_bias=True,
|
||||
processor=WanAttnProcessor2_0(),
|
||||
)
|
||||
self.norm2 = FP32LayerNorm(dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
|
||||
|
||||
# 3. Feed-forward
|
||||
self.ffn = FeedForward(dim, inner_dim=ffn_dim, activation_fn="gelu-approximate")
|
||||
self.norm3 = FP32LayerNorm(dim, eps, elementwise_affine=False)
|
||||
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
rotary_emb: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
|
||||
self.scale_shift_table + temb.float()
|
||||
).chunk(6, dim=1)
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) * (1 + scale_msa) + shift_msa).type_as(hidden_states)
|
||||
attn_output = self.attn1(hidden_states=norm_hidden_states, rotary_emb=rotary_emb)
|
||||
hidden_states = (hidden_states.float() + attn_output * gate_msa).type_as(hidden_states)
|
||||
|
||||
# 2. Cross-attention
|
||||
norm_hidden_states = self.norm2(hidden_states.float()).type_as(hidden_states)
|
||||
attn_output = self.attn2(hidden_states=norm_hidden_states, encoder_hidden_states=encoder_hidden_states)
|
||||
hidden_states = hidden_states + attn_output
|
||||
|
||||
# 3. Feed-forward
|
||||
norm_hidden_states = (self.norm3(hidden_states.float()) * (1 + c_scale_msa) + c_shift_msa).type_as(
|
||||
hidden_states
|
||||
)
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = (hidden_states.float() + ff_output.float() * c_gate_msa).type_as(hidden_states)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class WanTransformer3DModel(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin, CacheMixin):
|
||||
r"""
|
||||
A Transformer model for video-like data used in the Wan model.
|
||||
|
||||
Args:
|
||||
patch_size (`Tuple[int]`, defaults to `(1, 2, 2)`):
|
||||
3D patch dimensions for video embedding (t_patch, h_patch, w_patch).
|
||||
num_attention_heads (`int`, defaults to `40`):
|
||||
Fixed length for text embeddings.
|
||||
attention_head_dim (`int`, defaults to `128`):
|
||||
The number of channels in each head.
|
||||
in_channels (`int`, defaults to `16`):
|
||||
The number of channels in the input.
|
||||
out_channels (`int`, defaults to `16`):
|
||||
The number of channels in the output.
|
||||
text_dim (`int`, defaults to `512`):
|
||||
Input dimension for text embeddings.
|
||||
freq_dim (`int`, defaults to `256`):
|
||||
Dimension for sinusoidal time embeddings.
|
||||
ffn_dim (`int`, defaults to `13824`):
|
||||
Intermediate dimension in feed-forward network.
|
||||
num_layers (`int`, defaults to `40`):
|
||||
The number of layers of transformer blocks to use.
|
||||
window_size (`Tuple[int]`, defaults to `(-1, -1)`):
|
||||
Window size for local attention (-1 indicates global attention).
|
||||
cross_attn_norm (`bool`, defaults to `True`):
|
||||
Enable cross-attention normalization.
|
||||
qk_norm (`bool`, defaults to `True`):
|
||||
Enable query/key normalization.
|
||||
eps (`float`, defaults to `1e-6`):
|
||||
Epsilon value for normalization layers.
|
||||
add_img_emb (`bool`, defaults to `False`):
|
||||
Whether to use img_emb.
|
||||
added_kv_proj_dim (`int`, *optional*, defaults to `None`):
|
||||
The number of channels to use for the added key and value projections. If `None`, no projection is used.
|
||||
"""
|
||||
|
||||
_supports_gradient_checkpointing = True
|
||||
_skip_layerwise_casting_patterns = ["patch_embedding", "condition_embedder", "norm"]
|
||||
_no_split_modules = ["WanTransformerBlock"]
|
||||
_keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"]
|
||||
_keys_to_ignore_on_load_unexpected = ["norm_added_q"]
|
||||
|
||||
@register_to_config
|
||||
def __init__(
|
||||
self,
|
||||
patch_size: Tuple[int] = (1, 2, 2),
|
||||
num_attention_heads: int = 40,
|
||||
attention_head_dim: int = 128,
|
||||
in_channels: int = 16,
|
||||
out_channels: int = 16,
|
||||
text_dim: int = 4096,
|
||||
freq_dim: int = 256,
|
||||
ffn_dim: int = 13824,
|
||||
num_layers: int = 40,
|
||||
cross_attn_norm: bool = True,
|
||||
qk_norm: Optional[str] = "rms_norm_across_heads",
|
||||
eps: float = 1e-6,
|
||||
image_dim: Optional[int] = None,
|
||||
added_kv_proj_dim: Optional[int] = None,
|
||||
rope_max_seq_len: int = 1024,
|
||||
pos_embed_seq_len: Optional[int] = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
|
||||
inner_dim = num_attention_heads * attention_head_dim
|
||||
out_channels = out_channels or in_channels
|
||||
|
||||
# 1. Patch & position embedding
|
||||
self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len)
|
||||
self.patch_embedding = nn.Conv3d(in_channels, inner_dim, kernel_size=patch_size, stride=patch_size)
|
||||
|
||||
# 2. Condition embeddings
|
||||
# image_embedding_dim=1280 for I2V model
|
||||
self.condition_embedder = WanTimeTextImageEmbedding(
|
||||
dim=inner_dim,
|
||||
time_freq_dim=freq_dim,
|
||||
time_proj_dim=inner_dim * 6,
|
||||
text_embed_dim=text_dim,
|
||||
image_embed_dim=image_dim,
|
||||
pos_embed_seq_len=pos_embed_seq_len,
|
||||
)
|
||||
|
||||
# 3. Transformer blocks
|
||||
self.blocks = nn.ModuleList(
|
||||
[
|
||||
WanTransformerBlock(
|
||||
inner_dim, ffn_dim, num_attention_heads, qk_norm, cross_attn_norm, eps, added_kv_proj_dim
|
||||
)
|
||||
for _ in range(num_layers)
|
||||
]
|
||||
)
|
||||
|
||||
# 4. Output norm & projection
|
||||
self.norm_out = FP32LayerNorm(inner_dim, eps, elementwise_affine=False)
|
||||
self.proj_out = nn.Linear(inner_dim, out_channels * math.prod(patch_size))
|
||||
self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
timestep: torch.LongTensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
encoder_hidden_states_image: Optional[torch.Tensor] = None,
|
||||
return_dict: bool = True,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
|
||||
if attention_kwargs is not None:
|
||||
attention_kwargs = attention_kwargs.copy()
|
||||
lora_scale = attention_kwargs.pop("scale", 1.0)
|
||||
else:
|
||||
lora_scale = 1.0
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# weight the lora layers by setting `lora_scale` for each PEFT layer
|
||||
scale_lora_layers(self, lora_scale)
|
||||
else:
|
||||
if attention_kwargs is not None and attention_kwargs.get("scale", None) is not None:
|
||||
logger.warning(
|
||||
"Passing `scale` via `attention_kwargs` when not using the PEFT backend is ineffective."
|
||||
)
|
||||
|
||||
batch_size, num_channels, num_frames, height, width = hidden_states.shape
|
||||
p_t, p_h, p_w = self.config.patch_size
|
||||
post_patch_num_frames = num_frames // p_t
|
||||
post_patch_height = height // p_h
|
||||
post_patch_width = width // p_w
|
||||
|
||||
rotary_emb = self.rope(hidden_states)
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image
|
||||
)
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, -1))
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
|
||||
# 4. Transformer blocks
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block, hidden_states, encoder_hidden_states, timestep_proj, rotary_emb
|
||||
)
|
||||
else:
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states, timestep_proj, rotary_emb)
|
||||
|
||||
# 5. Output norm, projection & unpatchify
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
|
||||
|
||||
# Move the shift and scale tensors to the same device as hidden_states.
|
||||
# When using multi-GPU inference via accelerate these will be on the
|
||||
# first device rather than the last device, which hidden_states ends up
|
||||
# on.
|
||||
shift = shift.to(hidden_states.device)
|
||||
scale = scale.to(hidden_states.device)
|
||||
|
||||
hidden_states = (self.norm_out(hidden_states.float()) * (1 + scale) + shift).type_as(hidden_states)
|
||||
hidden_states = self.proj_out(hidden_states)
|
||||
|
||||
hidden_states = hidden_states.reshape(
|
||||
batch_size, post_patch_num_frames, post_patch_height, post_patch_width, p_t, p_h, p_w, -1
|
||||
)
|
||||
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
|
||||
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
|
||||
|
||||
if USE_PEFT_BACKEND:
|
||||
# remove `lora_scale` from each PEFT layer
|
||||
unscale_lora_layers(self, lora_scale)
|
||||
|
||||
if not return_dict:
|
||||
return (output,)
|
||||
|
||||
return Transformer2DModelOutput(sample=output)
|
||||
@@ -0,0 +1,609 @@
|
||||
# Copyright 2025 The Wan Team and The HuggingFace Team. All rights reserved.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import html
|
||||
from typing import Any, Callable, Dict, List, Optional, Union
|
||||
|
||||
import regex as re
|
||||
import torch
|
||||
from transformers import AutoTokenizer, UMT5EncoderModel
|
||||
|
||||
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback
|
||||
from diffusers.loaders import WanLoraLoaderMixin
|
||||
from diffusers.models import AutoencoderKLWan, WanTransformer3DModel
|
||||
from diffusers.schedulers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import is_ftfy_available, is_torch_xla_available, logging, replace_example_docstring
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
from diffusers.pipelines.pipeline_utils import DiffusionPipeline
|
||||
from diffusers.pipelines.wan.pipeline_output import WanPipelineOutput
|
||||
from einops import rearrange
|
||||
from transformers import UMT5EncoderModel, T5TokenizerFast
|
||||
|
||||
from fastvideo.models.mochi_hf.modeling_wan import WanTransformer3DModel
|
||||
from fastvideo.utils.communications import all_gather
|
||||
from fastvideo.utils.parallel_states import get_sequence_parallel_state, nccl_info
|
||||
|
||||
if is_torch_xla_available():
|
||||
import torch_xla.core.xla_model as xm
|
||||
|
||||
XLA_AVAILABLE = True
|
||||
else:
|
||||
XLA_AVAILABLE = False
|
||||
|
||||
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
|
||||
|
||||
if is_ftfy_available():
|
||||
import ftfy
|
||||
|
||||
|
||||
EXAMPLE_DOC_STRING = """
|
||||
Examples:
|
||||
```python
|
||||
>>> import torch
|
||||
>>> from diffusers.utils import export_to_video
|
||||
>>> from diffusers import AutoencoderKLWan, WanPipeline
|
||||
>>> from diffusers.schedulers.scheduling_unipc_multistep import UniPCMultistepScheduler
|
||||
|
||||
>>> # Available models: Wan-AI/Wan2.1-T2V-14B-Diffusers, Wan-AI/Wan2.1-T2V-1.3B-Diffusers
|
||||
>>> model_id = "Wan-AI/Wan2.1-T2V-14B-Diffusers"
|
||||
>>> vae = AutoencoderKLWan.from_pretrained(model_id, subfolder="vae", torch_dtype=torch.float32)
|
||||
>>> pipe = WanPipeline.from_pretrained(model_id, vae=vae, torch_dtype=torch.bfloat16)
|
||||
>>> flow_shift = 5.0 # 5.0 for 720P, 3.0 for 480P
|
||||
>>> pipe.scheduler = UniPCMultistepScheduler.from_config(pipe.scheduler.config, flow_shift=flow_shift)
|
||||
>>> pipe.to("cuda")
|
||||
|
||||
>>> prompt = "A cat and a dog baking a cake together in a kitchen. The cat is carefully measuring flour, while the dog is stirring the batter with a wooden spoon. The kitchen is cozy, with sunlight streaming through the window."
|
||||
>>> 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"
|
||||
|
||||
>>> output = pipe(
|
||||
... prompt=prompt,
|
||||
... negative_prompt=negative_prompt,
|
||||
... height=720,
|
||||
... width=1280,
|
||||
... num_frames=81,
|
||||
... guidance_scale=5.0,
|
||||
... ).frames[0]
|
||||
>>> export_to_video(output, "output.mp4", fps=16)
|
||||
```
|
||||
"""
|
||||
|
||||
|
||||
def basic_clean(text):
|
||||
text = ftfy.fix_text(text)
|
||||
text = html.unescape(html.unescape(text))
|
||||
return text.strip()
|
||||
|
||||
|
||||
def whitespace_clean(text):
|
||||
text = re.sub(r"\s+", " ", text)
|
||||
text = text.strip()
|
||||
return text
|
||||
|
||||
|
||||
def prompt_clean(text):
|
||||
text = whitespace_clean(basic_clean(text))
|
||||
return text
|
||||
|
||||
|
||||
class WanPipeline(DiffusionPipeline, WanLoraLoaderMixin):
|
||||
r"""
|
||||
Pipeline for text-to-video generation using Wan.
|
||||
|
||||
This model inherits from [`DiffusionPipeline`]. Check the superclass documentation for the generic methods
|
||||
implemented for all pipelines (downloading, saving, running on a particular device, etc.).
|
||||
|
||||
Args:
|
||||
tokenizer ([`T5Tokenizer`]):
|
||||
Tokenizer from [T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5Tokenizer),
|
||||
specifically the [google/umt5-xxl](https://huggingface.co/google/umt5-xxl) variant.
|
||||
text_encoder ([`T5EncoderModel`]):
|
||||
[T5](https://huggingface.co/docs/transformers/en/model_doc/t5#transformers.T5EncoderModel), specifically
|
||||
the [google/umt5-xxl](https://huggingface.co/google/umt5-xxl) variant.
|
||||
transformer ([`WanTransformer3DModel`]):
|
||||
Conditional Transformer to denoise the input latents.
|
||||
scheduler ([`UniPCMultistepScheduler`]):
|
||||
A scheduler to be used in combination with `transformer` to denoise the encoded image latents.
|
||||
vae ([`AutoencoderKLWan`]):
|
||||
Variational Auto-Encoder (VAE) Model to encode and decode videos to and from latent representations.
|
||||
"""
|
||||
|
||||
model_cpu_offload_seq = "text_encoder->transformer->vae"
|
||||
_callback_tensor_inputs = ["latents", "prompt_embeds", "negative_prompt_embeds"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
tokenizer: AutoTokenizer,
|
||||
text_encoder: UMT5EncoderModel,
|
||||
transformer: WanTransformer3DModel,
|
||||
vae: AutoencoderKLWan,
|
||||
scheduler: FlowMatchEulerDiscreteScheduler,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.register_modules(
|
||||
vae=vae,
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
transformer=transformer,
|
||||
scheduler=scheduler,
|
||||
)
|
||||
|
||||
self.vae_scale_factor_temporal = 2 ** sum(self.vae.temperal_downsample) if getattr(self, "vae", None) else 4
|
||||
self.vae_scale_factor_spatial = 2 ** len(self.vae.temperal_downsample) if getattr(self, "vae", None) else 8
|
||||
self.video_processor = VideoProcessor(vae_scale_factor=self.vae_scale_factor_spatial)
|
||||
|
||||
def _get_t5_prompt_embeds(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
num_videos_per_prompt: int = 1,
|
||||
max_sequence_length: int = 226,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
device = device or self._execution_device
|
||||
dtype = dtype or self.text_encoder.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
prompt = [prompt_clean(u) for u in prompt]
|
||||
batch_size = len(prompt)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=max_sequence_length,
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_attention_mask=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids, mask = text_inputs.input_ids, text_inputs.attention_mask
|
||||
seq_lens = mask.gt(0).sum(dim=1).long()
|
||||
|
||||
prompt_embeds = self.text_encoder(text_input_ids.to(device), mask.to(device)).last_hidden_state
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
prompt_embeds = [u[:v] for u, v in zip(prompt_embeds, seq_lens)]
|
||||
prompt_embeds = torch.stack(
|
||||
[torch.cat([u, u.new_zeros(max_sequence_length - u.size(0), u.size(1))]) for u in prompt_embeds], dim=0
|
||||
)
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
_, seq_len, _ = prompt_embeds.shape
|
||||
prompt_embeds = prompt_embeds.repeat(1, num_videos_per_prompt, 1)
|
||||
prompt_embeds = prompt_embeds.view(batch_size * num_videos_per_prompt, seq_len, -1)
|
||||
|
||||
return prompt_embeds
|
||||
|
||||
def encode_prompt(
|
||||
self,
|
||||
prompt: Union[str, List[str]],
|
||||
negative_prompt: Optional[Union[str, List[str]]] = None,
|
||||
do_classifier_free_guidance: bool = True,
|
||||
num_videos_per_prompt: int = 1,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
max_sequence_length: int = 226,
|
||||
device: Optional[torch.device] = None,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
):
|
||||
r"""
|
||||
Encodes the prompt into text encoder hidden states.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
prompt to be encoded
|
||||
negative_prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts not to guide the image generation. If not defined, one has to pass
|
||||
`negative_prompt_embeds` instead. Ignored when not using guidance (i.e., ignored if `guidance_scale` is
|
||||
less than `1`).
|
||||
do_classifier_free_guidance (`bool`, *optional*, defaults to `True`):
|
||||
Whether to use classifier free guidance or not.
|
||||
num_videos_per_prompt (`int`, *optional*, defaults to 1):
|
||||
Number of videos that should be generated per prompt. torch device to place the resulting embeddings on
|
||||
prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt weighting. If not
|
||||
provided, text embeddings will be generated from `prompt` input argument.
|
||||
negative_prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated negative text embeddings. Can be used to easily tweak text inputs, *e.g.* prompt
|
||||
weighting. If not provided, negative_prompt_embeds will be generated from `negative_prompt` input
|
||||
argument.
|
||||
device: (`torch.device`, *optional*):
|
||||
torch device
|
||||
dtype: (`torch.dtype`, *optional*):
|
||||
torch dtype
|
||||
"""
|
||||
device = device or self._execution_device
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
if prompt is not None:
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
if prompt_embeds is None:
|
||||
prompt_embeds = self._get_t5_prompt_embeds(
|
||||
prompt=prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
if do_classifier_free_guidance and negative_prompt_embeds is None:
|
||||
negative_prompt = negative_prompt or ""
|
||||
negative_prompt = batch_size * [negative_prompt] if isinstance(negative_prompt, str) else negative_prompt
|
||||
|
||||
if prompt is not None and type(prompt) is not type(negative_prompt):
|
||||
raise TypeError(
|
||||
f"`negative_prompt` should be the same type to `prompt`, but got {type(negative_prompt)} !="
|
||||
f" {type(prompt)}."
|
||||
)
|
||||
elif batch_size != len(negative_prompt):
|
||||
raise ValueError(
|
||||
f"`negative_prompt`: {negative_prompt} has batch size {len(negative_prompt)}, but `prompt`:"
|
||||
f" {prompt} has batch size {batch_size}. Please make sure that passed `negative_prompt` matches"
|
||||
" the batch size of `prompt`."
|
||||
)
|
||||
|
||||
negative_prompt_embeds = self._get_t5_prompt_embeds(
|
||||
prompt=negative_prompt,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
|
||||
return prompt_embeds, negative_prompt_embeds
|
||||
|
||||
def check_inputs(
|
||||
self,
|
||||
prompt,
|
||||
negative_prompt,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds=None,
|
||||
negative_prompt_embeds=None,
|
||||
callback_on_step_end_tensor_inputs=None,
|
||||
):
|
||||
if height % 16 != 0 or width % 16 != 0:
|
||||
raise ValueError(f"`height` and `width` have to be divisible by 16 but are {height} and {width}.")
|
||||
|
||||
if callback_on_step_end_tensor_inputs is not None and not all(
|
||||
k in self._callback_tensor_inputs for k in callback_on_step_end_tensor_inputs
|
||||
):
|
||||
raise ValueError(
|
||||
f"`callback_on_step_end_tensor_inputs` has to be in {self._callback_tensor_inputs}, but found {[k for k in callback_on_step_end_tensor_inputs if k not in self._callback_tensor_inputs]}"
|
||||
)
|
||||
|
||||
if prompt is not None and prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `prompt`: {prompt} and `prompt_embeds`: {prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two."
|
||||
)
|
||||
elif negative_prompt is not None and negative_prompt_embeds is not None:
|
||||
raise ValueError(
|
||||
f"Cannot forward both `negative_prompt`: {negative_prompt} and `negative_prompt_embeds`: {negative_prompt_embeds}. Please make sure to"
|
||||
" only forward one of the two."
|
||||
)
|
||||
elif prompt is None and prompt_embeds is None:
|
||||
raise ValueError(
|
||||
"Provide either `prompt` or `prompt_embeds`. Cannot leave both `prompt` and `prompt_embeds` undefined."
|
||||
)
|
||||
elif prompt is not None and (not isinstance(prompt, str) and not isinstance(prompt, list)):
|
||||
raise ValueError(f"`prompt` has to be of type `str` or `list` but is {type(prompt)}")
|
||||
elif negative_prompt is not None and (
|
||||
not isinstance(negative_prompt, str) and not isinstance(negative_prompt, list)
|
||||
):
|
||||
raise ValueError(f"`negative_prompt` has to be of type `str` or `list` but is {type(negative_prompt)}")
|
||||
|
||||
def prepare_latents(
|
||||
self,
|
||||
batch_size: int,
|
||||
num_channels_latents: int = 16,
|
||||
height: int = 480,
|
||||
width: int = 832,
|
||||
num_frames: int = 81,
|
||||
dtype: Optional[torch.dtype] = None,
|
||||
device: Optional[torch.device] = None,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
if latents is not None:
|
||||
return latents.to(device=device, dtype=dtype)
|
||||
|
||||
num_latent_frames = (num_frames - 1) // self.vae_scale_factor_temporal + 1
|
||||
shape = (
|
||||
batch_size,
|
||||
num_channels_latents,
|
||||
num_latent_frames,
|
||||
int(height) // self.vae_scale_factor_spatial,
|
||||
int(width) // self.vae_scale_factor_spatial,
|
||||
)
|
||||
if isinstance(generator, list) and len(generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"You have passed a list of generators of length {len(generator)}, but requested an effective batch"
|
||||
f" size of {batch_size}. Make sure the batch size matches the length of the generators."
|
||||
)
|
||||
|
||||
latents = randn_tensor(shape, generator=generator, device=device, dtype=dtype)
|
||||
return latents
|
||||
|
||||
@property
|
||||
def guidance_scale(self):
|
||||
return self._guidance_scale
|
||||
|
||||
@property
|
||||
def do_classifier_free_guidance(self):
|
||||
return self._guidance_scale > 1.0
|
||||
|
||||
@property
|
||||
def num_timesteps(self):
|
||||
return self._num_timesteps
|
||||
|
||||
@property
|
||||
def current_timestep(self):
|
||||
return self._current_timestep
|
||||
|
||||
@property
|
||||
def interrupt(self):
|
||||
return self._interrupt
|
||||
|
||||
@property
|
||||
def attention_kwargs(self):
|
||||
return self._attention_kwargs
|
||||
|
||||
@torch.no_grad()
|
||||
@replace_example_docstring(EXAMPLE_DOC_STRING)
|
||||
def __call__(
|
||||
self,
|
||||
prompt: Union[str, List[str]] = None,
|
||||
negative_prompt: Union[str, List[str]] = None,
|
||||
height: int = 480,
|
||||
width: int = 832,
|
||||
num_frames: int = 81,
|
||||
num_inference_steps: int = 50,
|
||||
guidance_scale: float = 5.0,
|
||||
num_videos_per_prompt: Optional[int] = 1,
|
||||
generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None,
|
||||
latents: Optional[torch.Tensor] = None,
|
||||
prompt_embeds: Optional[torch.Tensor] = None,
|
||||
negative_prompt_embeds: Optional[torch.Tensor] = None,
|
||||
output_type: Optional[str] = "np",
|
||||
return_dict: bool = True,
|
||||
attention_kwargs: Optional[Dict[str, Any]] = None,
|
||||
callback_on_step_end: Optional[
|
||||
Union[Callable[[int, int, Dict], None], PipelineCallback, MultiPipelineCallbacks]
|
||||
] = None,
|
||||
callback_on_step_end_tensor_inputs: List[str] = ["latents"],
|
||||
max_sequence_length: int = 512,
|
||||
):
|
||||
r"""
|
||||
The call function to the pipeline for generation.
|
||||
|
||||
Args:
|
||||
prompt (`str` or `List[str]`, *optional*):
|
||||
The prompt or prompts to guide the image generation. If not defined, one has to pass `prompt_embeds`.
|
||||
instead.
|
||||
height (`int`, defaults to `480`):
|
||||
The height in pixels of the generated image.
|
||||
width (`int`, defaults to `832`):
|
||||
The width in pixels of the generated image.
|
||||
num_frames (`int`, defaults to `81`):
|
||||
The number of frames in the generated video.
|
||||
num_inference_steps (`int`, defaults to `50`):
|
||||
The number of denoising steps. More denoising steps usually lead to a higher quality image at the
|
||||
expense of slower inference.
|
||||
guidance_scale (`float`, defaults to `5.0`):
|
||||
Guidance scale as defined in [Classifier-Free Diffusion
|
||||
Guidance](https://huggingface.co/papers/2207.12598). `guidance_scale` is defined as `w` of equation 2.
|
||||
of [Imagen Paper](https://huggingface.co/papers/2205.11487). Guidance scale is enabled by setting
|
||||
`guidance_scale > 1`. Higher guidance scale encourages to generate images that are closely linked to
|
||||
the text `prompt`, usually at the expense of lower image quality.
|
||||
num_videos_per_prompt (`int`, *optional*, defaults to 1):
|
||||
The number of images to generate per prompt.
|
||||
generator (`torch.Generator` or `List[torch.Generator]`, *optional*):
|
||||
A [`torch.Generator`](https://pytorch.org/docs/stable/generated/torch.Generator.html) to make
|
||||
generation deterministic.
|
||||
latents (`torch.Tensor`, *optional*):
|
||||
Pre-generated noisy latents sampled from a Gaussian distribution, to be used as inputs for image
|
||||
generation. Can be used to tweak the same generation with different prompts. If not provided, a latents
|
||||
tensor is generated by sampling using the supplied random `generator`.
|
||||
prompt_embeds (`torch.Tensor`, *optional*):
|
||||
Pre-generated text embeddings. Can be used to easily tweak text inputs (prompt weighting). If not
|
||||
provided, text embeddings are generated from the `prompt` input argument.
|
||||
output_type (`str`, *optional*, defaults to `"np"`):
|
||||
The output format of the generated image. Choose between `PIL.Image` or `np.array`.
|
||||
return_dict (`bool`, *optional*, defaults to `True`):
|
||||
Whether or not to return a [`WanPipelineOutput`] instead of a plain tuple.
|
||||
attention_kwargs (`dict`, *optional*):
|
||||
A kwargs dictionary that if specified is passed along to the `AttentionProcessor` as defined under
|
||||
`self.processor` in
|
||||
[diffusers.models.attention_processor](https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/attention_processor.py).
|
||||
callback_on_step_end (`Callable`, `PipelineCallback`, `MultiPipelineCallbacks`, *optional*):
|
||||
A function or a subclass of `PipelineCallback` or `MultiPipelineCallbacks` that is called at the end of
|
||||
each denoising step during the inference. with the following arguments: `callback_on_step_end(self:
|
||||
DiffusionPipeline, step: int, timestep: int, callback_kwargs: Dict)`. `callback_kwargs` will include a
|
||||
list of all tensors as specified by `callback_on_step_end_tensor_inputs`.
|
||||
callback_on_step_end_tensor_inputs (`List`, *optional*):
|
||||
The list of tensor inputs for the `callback_on_step_end` function. The tensors specified in the list
|
||||
will be passed as `callback_kwargs` argument. You will only be able to include variables listed in the
|
||||
`._callback_tensor_inputs` attribute of your pipeline class.
|
||||
autocast_dtype (`torch.dtype`, *optional*, defaults to `torch.bfloat16`):
|
||||
The dtype to use for the torch.amp.autocast.
|
||||
|
||||
Examples:
|
||||
|
||||
Returns:
|
||||
[`~WanPipelineOutput`] or `tuple`:
|
||||
If `return_dict` is `True`, [`WanPipelineOutput`] is returned, otherwise a `tuple` is returned where
|
||||
the first element is a list with the generated images and the second element is a list of `bool`s
|
||||
indicating whether the corresponding generated image contains "not-safe-for-work" (nsfw) content.
|
||||
"""
|
||||
|
||||
if isinstance(callback_on_step_end, (PipelineCallback, MultiPipelineCallbacks)):
|
||||
callback_on_step_end_tensor_inputs = callback_on_step_end.tensor_inputs
|
||||
|
||||
# 1. Check inputs. Raise error if not correct
|
||||
self.check_inputs(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
height,
|
||||
width,
|
||||
prompt_embeds,
|
||||
negative_prompt_embeds,
|
||||
callback_on_step_end_tensor_inputs,
|
||||
)
|
||||
|
||||
if num_frames % self.vae_scale_factor_temporal != 1:
|
||||
logger.warning(
|
||||
f"`num_frames - 1` has to be divisible by {self.vae_scale_factor_temporal}. Rounding to the nearest number."
|
||||
)
|
||||
num_frames = num_frames // self.vae_scale_factor_temporal * self.vae_scale_factor_temporal + 1
|
||||
num_frames = max(num_frames, 1)
|
||||
|
||||
self._guidance_scale = guidance_scale
|
||||
self._attention_kwargs = attention_kwargs
|
||||
self._current_timestep = None
|
||||
self._interrupt = False
|
||||
|
||||
device = self._execution_device
|
||||
|
||||
# 2. Define call parameters
|
||||
if prompt is not None and isinstance(prompt, str):
|
||||
batch_size = 1
|
||||
elif prompt is not None and isinstance(prompt, list):
|
||||
batch_size = len(prompt)
|
||||
else:
|
||||
batch_size = prompt_embeds.shape[0]
|
||||
|
||||
# 3. Encode input prompt
|
||||
prompt_embeds, negative_prompt_embeds = self.encode_prompt(
|
||||
prompt=prompt,
|
||||
negative_prompt=negative_prompt,
|
||||
do_classifier_free_guidance=self.do_classifier_free_guidance,
|
||||
num_videos_per_prompt=num_videos_per_prompt,
|
||||
prompt_embeds=prompt_embeds,
|
||||
negative_prompt_embeds=negative_prompt_embeds,
|
||||
max_sequence_length=max_sequence_length,
|
||||
device=device,
|
||||
)
|
||||
|
||||
transformer_dtype = self.transformer.dtype
|
||||
prompt_embeds = prompt_embeds.to(transformer_dtype)
|
||||
if negative_prompt_embeds is not None:
|
||||
negative_prompt_embeds = negative_prompt_embeds.to(transformer_dtype)
|
||||
|
||||
# 4. Prepare timesteps
|
||||
self.scheduler.set_timesteps(num_inference_steps, device=device)
|
||||
timesteps = self.scheduler.timesteps
|
||||
|
||||
# 5. Prepare latent variables
|
||||
num_channels_latents = self.transformer.config.in_channels
|
||||
latents = self.prepare_latents(
|
||||
batch_size * num_videos_per_prompt,
|
||||
num_channels_latents,
|
||||
height,
|
||||
width,
|
||||
num_frames,
|
||||
torch.float32,
|
||||
device,
|
||||
generator,
|
||||
latents,
|
||||
)
|
||||
|
||||
world_size, rank = nccl_info.sp_size, nccl_info.rank_within_group
|
||||
if get_sequence_parallel_state():
|
||||
latents = rearrange(latents, "b t (n s) h w -> b t n s h w", n=world_size).contiguous()
|
||||
latents = latents[:, :, rank, :, :, :]
|
||||
|
||||
# 6. Denoising loop
|
||||
num_warmup_steps = len(timesteps) - num_inference_steps * self.scheduler.order
|
||||
self._num_timesteps = len(timesteps)
|
||||
|
||||
self._progress_bar_config = {"disable": nccl_info.rank_within_group != 0}
|
||||
with self.progress_bar(total=num_inference_steps) as progress_bar:
|
||||
for i, t in enumerate(timesteps):
|
||||
if self.interrupt:
|
||||
continue
|
||||
|
||||
self._current_timestep = t
|
||||
latent_model_input = latents.to(transformer_dtype)
|
||||
timestep = t.expand(latents.shape[0])
|
||||
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
if self.do_classifier_free_guidance:
|
||||
noise_uncond = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=negative_prompt_embeds,
|
||||
attention_kwargs=attention_kwargs,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
noise_pred = noise_uncond + guidance_scale * (noise_pred - noise_uncond)
|
||||
|
||||
# compute the previous noisy sample x_t -> x_t-1
|
||||
latents = self.scheduler.step(noise_pred, t, latents, return_dict=False)[0]
|
||||
|
||||
if callback_on_step_end is not None:
|
||||
callback_kwargs = {}
|
||||
for k in callback_on_step_end_tensor_inputs:
|
||||
callback_kwargs[k] = locals()[k]
|
||||
callback_outputs = callback_on_step_end(self, i, t, callback_kwargs)
|
||||
|
||||
latents = callback_outputs.pop("latents", latents)
|
||||
prompt_embeds = callback_outputs.pop("prompt_embeds", prompt_embeds)
|
||||
negative_prompt_embeds = callback_outputs.pop("negative_prompt_embeds", negative_prompt_embeds)
|
||||
|
||||
# call the callback, if provided
|
||||
if i == len(timesteps) - 1 or ((i + 1) > num_warmup_steps and (i + 1) % self.scheduler.order == 0):
|
||||
progress_bar.update()
|
||||
|
||||
if XLA_AVAILABLE:
|
||||
xm.mark_step()
|
||||
|
||||
if get_sequence_parallel_state():
|
||||
latents = all_gather(latents, dim=2)
|
||||
|
||||
self._current_timestep = None
|
||||
|
||||
if not output_type == "latent":
|
||||
latents = latents.to(self.vae.dtype)
|
||||
latents_mean = (
|
||||
torch.tensor(self.vae.config.latents_mean)
|
||||
.view(1, self.vae.config.z_dim, 1, 1, 1)
|
||||
.to(latents.device, latents.dtype)
|
||||
)
|
||||
latents_std = 1.0 / torch.tensor(self.vae.config.latents_std).view(1, self.vae.config.z_dim, 1, 1, 1).to(
|
||||
latents.device, latents.dtype
|
||||
)
|
||||
latents = latents / latents_std + latents_mean
|
||||
video = self.vae.decode(latents, return_dict=False)[0]
|
||||
video = self.video_processor.postprocess_video(video, output_type=output_type)
|
||||
else:
|
||||
video = latents
|
||||
|
||||
# Offload all models
|
||||
self.maybe_free_model_hooks()
|
||||
|
||||
if not return_dict:
|
||||
return (video,)
|
||||
|
||||
return WanPipelineOutput(frames=video)
|
||||
@@ -86,7 +86,7 @@ def inference(args):
|
||||
num_inference_steps=args.num_inference_steps,
|
||||
generator=generator,
|
||||
).frames
|
||||
if nccl_info.global_rank <= 0:
|
||||
if nccl_info.global_rank == 0:
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
suffix = prompt.split(".")[0]
|
||||
export_to_video(
|
||||
@@ -107,7 +107,7 @@ def inference(args):
|
||||
generator=generator,
|
||||
).frames
|
||||
|
||||
if nccl_info.global_rank <= 0:
|
||||
if nccl_info.global_rank == 0:
|
||||
export_to_video(videos[0], args.output_path + ".mp4", fps=24)
|
||||
|
||||
|
||||
|
||||
@@ -94,7 +94,7 @@ def main(args):
|
||||
guidance_scale=args.guidance_scale,
|
||||
generator=generator,
|
||||
).frames
|
||||
if nccl_info.global_rank <= 0:
|
||||
if nccl_info.global_rank == 0:
|
||||
os.makedirs(args.output_path, exist_ok=True)
|
||||
suffix = prompt.split(".")[0]
|
||||
export_to_video(
|
||||
@@ -116,7 +116,7 @@ def main(args):
|
||||
generator=generator,
|
||||
).frames
|
||||
|
||||
if nccl_info.global_rank <= 0:
|
||||
if nccl_info.global_rank == 0:
|
||||
export_to_video(videos[0], args.output_path + ".mp4", fps=30)
|
||||
|
||||
|
||||
|
||||
+5
-11
@@ -20,7 +20,7 @@ from torch.utils.data.distributed import DistributedSampler
|
||||
from tqdm.auto import tqdm
|
||||
|
||||
from fastvideo.dataset.latent_datasets import (LatentDataset, latent_collate_function)
|
||||
from fastvideo.models.mochi_hf.mochi_latents_utils import normalize_dit_input
|
||||
from fastvideo.utils.latents_utils import normalize_dit_input
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
from fastvideo.models.hunyuan_hf.pipeline_hunyuan import HunyuanVideoPipeline
|
||||
|
||||
@@ -185,7 +185,7 @@ def main(args):
|
||||
noise_random_generator = None
|
||||
|
||||
# Handle the repository creation
|
||||
if rank <= 0 and args.output_dir is not None:
|
||||
if rank == 0 and args.output_dir is not None:
|
||||
os.makedirs(args.output_dir, exist_ok=True)
|
||||
|
||||
# For mixed precision training we cast all non-trainable weights to half-precision
|
||||
@@ -316,7 +316,7 @@ def main(args):
|
||||
len(train_dataloader) / args.gradient_accumulation_steps * args.sp_size / args.train_sp_batch_size)
|
||||
args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch)
|
||||
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
project = args.tracker_project_name or "fastvideo"
|
||||
wandb.init(project=project, config=args)
|
||||
|
||||
@@ -393,7 +393,7 @@ def main(args):
|
||||
"grad_norm": grad_norm,
|
||||
})
|
||||
progress_bar.update(1)
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
wandb.log(
|
||||
{
|
||||
"train_loss": loss,
|
||||
@@ -520,13 +520,7 @@ if __name__ == "__main__":
|
||||
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
|
||||
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--logging_dir",
|
||||
type=str,
|
||||
default="logs",
|
||||
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
|
||||
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
|
||||
)
|
||||
|
||||
|
||||
# optimizer & scheduler & Training
|
||||
parser.add_argument("--num_train_epochs", type=int, default=100)
|
||||
|
||||
@@ -32,7 +32,7 @@ def save_checkpoint_optimizer(model, optimizer, rank, output_dir, step, discrimi
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
if rank <= 0 and not discriminator:
|
||||
if rank == 0 and not discriminator:
|
||||
weight_path = os.path.join(save_dir, "diffusion_pytorch_model.safetensors")
|
||||
save_file(cpu_state, weight_path)
|
||||
config_dict = dict(model.config)
|
||||
@@ -60,7 +60,7 @@ def save_checkpoint(transformer, rank, output_dir, step):
|
||||
):
|
||||
cpu_state = transformer.state_dict()
|
||||
# todo move to get_state_dict
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
save_dir = os.path.join(output_dir, f"checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
@@ -98,7 +98,7 @@ def save_checkpoint_generator_discriminator(
|
||||
hf_weight_dir = os.path.join(save_dir, "hf_weights")
|
||||
os.makedirs(hf_weight_dir, exist_ok=True)
|
||||
# save using safetensors
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
config_dict = dict(model.config)
|
||||
config_path = os.path.join(hf_weight_dir, "config.json")
|
||||
# save dict as json
|
||||
@@ -139,7 +139,7 @@ def save_checkpoint_generator_discriminator(
|
||||
optim_state = FSDP.optim_state_dict(discriminator, discriminator_optimizer)
|
||||
model_state = discriminator.state_dict()
|
||||
state_dict = {"optimizer": optim_state, "model": model_state}
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
discriminator_fsdp_state_fil = os.path.join(discriminator_fsdp_state_dir, "discriminator_state.pt")
|
||||
torch.save(state_dict, discriminator_fsdp_state_fil)
|
||||
|
||||
@@ -178,7 +178,7 @@ def load_full_state_model(model, optimizer, checkpoint_file, rank):
|
||||
):
|
||||
discriminator_state = torch.load(checkpoint_file)
|
||||
model_state = discriminator_state["model"]
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
optim_state = discriminator_state["optimizer"]
|
||||
else:
|
||||
optim_state = None
|
||||
@@ -241,7 +241,7 @@ def save_lora_checkpoint(transformer, optimizer, rank, output_dir, step, pipelin
|
||||
optimizer,
|
||||
)
|
||||
|
||||
if rank <= 0:
|
||||
if rank == 0:
|
||||
save_dir = os.path.join(output_dir, f"lora-checkpoint-{step}")
|
||||
os.makedirs(save_dir, exist_ok=True)
|
||||
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
import torch
|
||||
|
||||
mochi_latents_mean = torch.tensor([
|
||||
-0.06730895953510081,
|
||||
-0.038011381506090416,
|
||||
-0.07477820912866141,
|
||||
-0.05565264470995561,
|
||||
0.012767231469026969,
|
||||
-0.04703542746246419,
|
||||
0.043896967884726704,
|
||||
-0.09346305707025976,
|
||||
-0.09918314763016893,
|
||||
-0.008729793427399178,
|
||||
-0.011931556316503654,
|
||||
-0.0321993391887285,
|
||||
]).view(1, 12, 1, 1, 1)
|
||||
mochi_latents_std = torch.tensor([
|
||||
0.9263795028493863,
|
||||
0.9248894543193766,
|
||||
0.9393059390890617,
|
||||
0.959253732819592,
|
||||
0.8244560132752793,
|
||||
0.917259975397747,
|
||||
0.9294154431013696,
|
||||
1.3720942357788521,
|
||||
0.881393668867029,
|
||||
0.9168315692124348,
|
||||
0.9185249279345552,
|
||||
0.9274757570805041,
|
||||
]).view(1, 12, 1, 1, 1)
|
||||
mochi_scaling_factor = 1.0
|
||||
|
||||
|
||||
wan_latents_mean = torch.tensor([
|
||||
-0.7571,
|
||||
-0.7089,
|
||||
-0.9113,
|
||||
0.1075,
|
||||
-0.1745,
|
||||
0.9653,
|
||||
-0.1517,
|
||||
1.5508,
|
||||
0.4134,
|
||||
-0.0715,
|
||||
0.5517,
|
||||
-0.3632,
|
||||
-0.1922,
|
||||
-0.9497,
|
||||
0.2503,
|
||||
-0.2921,
|
||||
]).view(1, 16, 1, 1, 1)
|
||||
wan_latents_std = torch.tensor([
|
||||
2.8184,
|
||||
1.4541,
|
||||
2.3275,
|
||||
2.6558,
|
||||
1.2196,
|
||||
1.7708,
|
||||
2.6052,
|
||||
2.0743,
|
||||
3.2687,
|
||||
2.1526,
|
||||
2.8652,
|
||||
1.5579,
|
||||
1.6382,
|
||||
1.1253,
|
||||
2.8251,
|
||||
1.916,
|
||||
]).view(1, 16, 1, 1, 1)
|
||||
|
||||
|
||||
def normalize_dit_input(model_type, latents):
|
||||
if model_type == "mochi":
|
||||
latents_mean = mochi_latents_mean.to(latents.device, latents.dtype)
|
||||
latents_std = mochi_latents_std.to(latents.device, latents.dtype)
|
||||
latents = (latents - latents_mean) / latents_std
|
||||
return latents
|
||||
elif model_type == "hunyuan_hf":
|
||||
return latents * 0.476986
|
||||
elif model_type == "hunyuan":
|
||||
return latents * 0.476986
|
||||
elif model_type == "wan":
|
||||
latents_mean = wan_latents_mean.to(latents.device, latents.dtype)
|
||||
latents_std = wan_latents_std.to(latents.device, latents.dtype)
|
||||
latents = (latents - latents_mean) / latents_std
|
||||
return latents
|
||||
else:
|
||||
raise NotImplementedError(f"model_type {model_type} not supported")
|
||||
+69
-2
@@ -3,9 +3,9 @@ from pathlib import Path
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from diffusers import AutoencoderKLHunyuanVideo, AutoencoderKLMochi
|
||||
from diffusers import AutoencoderKLHunyuanVideo, AutoencoderKLMochi, AutoencoderKLWan
|
||||
from torch import nn
|
||||
from transformers import AutoTokenizer, T5EncoderModel
|
||||
from transformers import AutoTokenizer, T5EncoderModel, UMT5EncoderModel
|
||||
|
||||
from fastvideo.models.hunyuan.modules.models import (HYVideoDiffusionTransformer, MMDoubleStreamBlock,
|
||||
MMSingleStreamBlock)
|
||||
@@ -14,6 +14,7 @@ from fastvideo.models.hunyuan.vae.autoencoder_kl_causal_3d import AutoencoderKLC
|
||||
from fastvideo.models.hunyuan_hf.modeling_hunyuan import (HunyuanVideoSingleTransformerBlock,
|
||||
HunyuanVideoTransformer3DModel, HunyuanVideoTransformerBlock)
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel, MochiTransformerBlock
|
||||
from fastvideo.models.wan_hf.modeling_wan import WanTransformer3DModel, WanTransformerBlock
|
||||
from fastvideo.utils.logging_ import main_print
|
||||
|
||||
hunyuan_config = {
|
||||
@@ -200,6 +201,48 @@ class MochiTextEncoderWrapper(nn.Module):
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
class WanTextEncoderWrapper(nn.Module):
|
||||
|
||||
def __init__(self, pretrained_model_name_or_path, device):
|
||||
super().__init__()
|
||||
self.text_encoder = UMT5EncoderModel.from_pretrained(os.path.join(pretrained_model_name_or_path,
|
||||
"text_encoder")).to(device)
|
||||
self.tokenizer = AutoTokenizer.from_pretrained(os.path.join(pretrained_model_name_or_path, "tokenizer"))
|
||||
self.max_sequence_length = 256
|
||||
|
||||
def encode_prompt(self, prompt):
|
||||
device = self.text_encoder.device
|
||||
dtype = self.text_encoder.dtype
|
||||
|
||||
prompt = [prompt] if isinstance(prompt, str) else prompt
|
||||
batch_size = len(prompt)
|
||||
|
||||
text_inputs = self.tokenizer(
|
||||
prompt,
|
||||
padding="max_length",
|
||||
max_length=self.max_sequence_length,
|
||||
truncation=True,
|
||||
add_special_tokens=True,
|
||||
return_tensors="pt",
|
||||
)
|
||||
text_input_ids = text_inputs.input_ids
|
||||
prompt_attention_mask = text_inputs.attention_mask
|
||||
prompt_attention_mask = prompt_attention_mask.bool().to(device)
|
||||
|
||||
untruncated_ids = self.tokenizer(prompt, padding="longest", return_tensors="pt").input_ids
|
||||
|
||||
if untruncated_ids.shape[-1] >= text_input_ids.shape[-1] and not torch.equal(text_input_ids, untruncated_ids):
|
||||
removed_text = self.tokenizer.batch_decode(untruncated_ids[:, self.max_sequence_length - 1:-1])
|
||||
main_print(f"Truncated text input: {prompt} to: {removed_text} for model input.")
|
||||
prompt_embeds = self.text_encoder(text_input_ids.to(device), attention_mask=prompt_attention_mask)[0]
|
||||
prompt_embeds = prompt_embeds.to(dtype=dtype, device=device)
|
||||
|
||||
# duplicate text embeddings for each generation per prompt, using mps friendly method
|
||||
_, seq_len, _ = prompt_embeds.shape
|
||||
prompt_embeds = prompt_embeds.view(batch_size, seq_len, -1)
|
||||
prompt_attention_mask = prompt_attention_mask.view(batch_size, -1)
|
||||
|
||||
return prompt_embeds, prompt_attention_mask
|
||||
|
||||
def load_hunyuan_state_dict(model, dit_model_name_or_path):
|
||||
load_key = "module"
|
||||
@@ -240,6 +283,20 @@ def load_transformer(
|
||||
torch_dtype=master_weight_type,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
elif model_type == "wan":
|
||||
if dit_model_name_or_path:
|
||||
transformer = WanTransformer3DModel.from_pretrained(
|
||||
dit_model_name_or_path,
|
||||
torch_dtype=master_weight_type,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
else:
|
||||
transformer = WanTransformer3DModel.from_pretrained(
|
||||
pretrained_model_name_or_path,
|
||||
subfolder="transformer",
|
||||
torch_dtype=master_weight_type,
|
||||
# torch_dtype=torch.bfloat16 if args.use_lora else torch.float32,
|
||||
)
|
||||
elif model_type == "hunyuan_hf":
|
||||
if dit_model_name_or_path:
|
||||
transformer = HunyuanVideoTransformer3DModel.from_pretrained(
|
||||
@@ -283,6 +340,12 @@ def load_vae(model_type, pretrained_model_name_or_path):
|
||||
torch_dtype=weight_dtype).to("cuda")
|
||||
autocast_type = torch.bfloat16
|
||||
fps = 24
|
||||
elif model_type == "wan":
|
||||
vae = AutoencoderKLWan.from_pretrained(pretrained_model_name_or_path,
|
||||
subfolder="vae",
|
||||
torch_dtype=weight_dtype).to("cuda")
|
||||
autocast_type = torch.bfloat16
|
||||
fps = 24
|
||||
elif model_type == "hunyuan":
|
||||
vae_precision = torch.float32
|
||||
vae_path = os.path.join(pretrained_model_name_or_path, "hunyuan-video-t2v-720p/vae")
|
||||
@@ -311,6 +374,8 @@ def load_vae(model_type, pretrained_model_name_or_path):
|
||||
def load_text_encoder(model_type, pretrained_model_name_or_path, device):
|
||||
if model_type == "mochi":
|
||||
text_encoder = MochiTextEncoderWrapper(pretrained_model_name_or_path, device)
|
||||
elif model_type == "wan":
|
||||
text_encoder = WanTextEncoderWrapper(pretrained_model_name_or_path, device)
|
||||
elif model_type == "hunyuan" or "hunyuan_hf":
|
||||
text_encoder = HunyuanTextEncoderWrapper(pretrained_model_name_or_path, device)
|
||||
else:
|
||||
@@ -322,6 +387,8 @@ def get_no_split_modules(transformer):
|
||||
# if of type MochiTransformer3DModel
|
||||
if isinstance(transformer, MochiTransformer3DModel):
|
||||
return (MochiTransformerBlock, )
|
||||
elif isinstance(transformer, WanTransformer3DModel):
|
||||
return (WanTransformerBlock, )
|
||||
elif isinstance(transformer, HunyuanVideoTransformer3DModel):
|
||||
return (HunyuanVideoSingleTransformerBlock, HunyuanVideoTransformerBlock)
|
||||
elif isinstance(transformer, HYVideoDiffusionTransformer):
|
||||
|
||||
@@ -129,13 +129,22 @@ def sample_validation_video(
|
||||
# broadcast to batch dimension in a way that's compatible with ONNX/Core ML
|
||||
timestep = t.expand(latent_model_input.shape[0])
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
noise_pred = transformer(
|
||||
hidden_states=latent_model_input,
|
||||
encoder_hidden_states=prompt_embeds,
|
||||
timestep=timestep,
|
||||
encoder_attention_mask=prompt_attention_mask,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
if model_type == "wan":
|
||||
pred_kwargs = {
|
||||
"hidden_states": latent_model_input,
|
||||
"encoder_hidden_states": prompt_embeds,
|
||||
"timestep":timestep,
|
||||
"return_dict":False,
|
||||
}
|
||||
else:
|
||||
pred_kwargs = {
|
||||
"hidden_states": latent_model_input,
|
||||
"encoder_hidden_states": prompt_embeds,
|
||||
"timestep":timestep,
|
||||
"encoder_attention_mask":prompt_attention_mask,
|
||||
"return_dict":False,
|
||||
}
|
||||
noise_pred = transformer(**pred_kwargs)[0]
|
||||
|
||||
# Mochi CFG + Sampling runs in FP32
|
||||
noise_pred = noise_pred.to(torch.float32)
|
||||
@@ -166,10 +175,12 @@ def sample_validation_video(
|
||||
# denormalize with the mean and std if available and not None
|
||||
has_latents_mean = (hasattr(vae.config, "latents_mean") and vae.config.latents_mean is not None)
|
||||
has_latents_std = (hasattr(vae.config, "latents_std") and vae.config.latents_std is not None)
|
||||
if model_type == "wan":
|
||||
vae.config.scaling_factor = 1
|
||||
if has_latents_mean and has_latents_std:
|
||||
latents_mean = (torch.tensor(vae.config.latents_mean).view(1, 12, 1, 1,
|
||||
latents_mean = (torch.tensor(vae.config.latents_mean).view(1, num_channels_latents, 1, 1,
|
||||
1).to(latents.device, latents.dtype))
|
||||
latents_std = (torch.tensor(vae.config.latents_std).view(1, 12, 1, 1, 1).to(latents.device, latents.dtype))
|
||||
latents_std = (torch.tensor(vae.config.latents_std).view(1, num_channels_latents, 1, 1, 1).to(latents.device, latents.dtype))
|
||||
latents = latents * latents_std / vae.config.scaling_factor + latents_mean
|
||||
else:
|
||||
latents = latents / vae.config.scaling_factor
|
||||
@@ -202,14 +213,15 @@ def log_validation(
|
||||
vae_spatial_scale_factor = 8
|
||||
vae_temporal_scale_factor = 6
|
||||
num_channels_latents = 12
|
||||
elif args.model_type == "hunyuan" or "hunyuan_hf":
|
||||
elif args.model_type == "hunyuan" or "hunyuan_hf" or "wan":
|
||||
vae_spatial_scale_factor = 8
|
||||
vae_temporal_scale_factor = 4
|
||||
num_channels_latents = 16
|
||||
else:
|
||||
raise ValueError(f"Model type {args.model_type} not supported")
|
||||
vae, autocast_type, fps = load_vae(args.model_type, args.pretrained_model_name_or_path)
|
||||
vae.enable_tiling()
|
||||
if args.model_type != "wan":
|
||||
vae.enable_tiling()
|
||||
if scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler(shift=shift)
|
||||
else:
|
||||
|
||||
@@ -0,0 +1,420 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import os
|
||||
from collections import defaultdict
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
import numpy as np
|
||||
|
||||
|
||||
def configure_sta(mode: str = 'STA_searching',
|
||||
layer_num: int = 40,
|
||||
time_step_num: int = 50,
|
||||
head_num: int = 40,
|
||||
**kwargs) -> List[List[List[Any]]]:
|
||||
"""
|
||||
Configure Sliding Tile Attention (STA) parameters based on the specified mode.
|
||||
|
||||
Parameters:
|
||||
----------
|
||||
mode : str
|
||||
The STA mode to use. Options are:
|
||||
- 'STA_searching': Generate a set of mask candidates for initial search
|
||||
- 'STA_tuning': Select best mask strategy based on previously saved results
|
||||
- 'STA_inference': Load and use a previously tuned mask strategy
|
||||
layer_num: int, number of layers
|
||||
time_step_num: int, number of timesteps
|
||||
head_num: int, number of heads
|
||||
|
||||
**kwargs : dict
|
||||
Mode-specific parameters:
|
||||
|
||||
For 'STA_searching':
|
||||
- mask_candidates: list of str, optional, mask candidates to use
|
||||
- mask_selected: list of int, optional, indices of selected masks
|
||||
|
||||
For 'STA_tuning':
|
||||
- mask_search_files_path: str, required, path to mask search results
|
||||
- mask_candidates: list of str, optional, mask candidates to use
|
||||
- mask_selected: list of int, optional, indices of selected masks
|
||||
- skip_time_steps: int, optional, number of time steps to use full attention (default 12)
|
||||
- save_dir: str, optional, directory to save mask strategy (default "mask_candidates")
|
||||
|
||||
For 'STA_inference':
|
||||
- load_path: str, optional, path to load mask strategy (default "mask_candidates/mask_strategy.json")
|
||||
"""
|
||||
valid_modes = [
|
||||
'STA_searching', 'STA_tuning', 'STA_inference', 'STA_tuning_cfg'
|
||||
]
|
||||
if mode not in valid_modes:
|
||||
raise ValueError(f"Mode must be one of {valid_modes}, got {mode}")
|
||||
|
||||
if mode == 'STA_searching':
|
||||
# Get parameters with defaults
|
||||
mask_candidates: Optional[List[str]] = kwargs.get('mask_candidates')
|
||||
if mask_candidates is None:
|
||||
raise ValueError(
|
||||
"mask_candidates is required for STA_searching mode")
|
||||
mask_selected: List[int] = kwargs.get('mask_selected',
|
||||
list(range(len(mask_candidates))))
|
||||
|
||||
# Parse selected masks
|
||||
selected_masks: List[List[int]] = []
|
||||
for index in mask_selected:
|
||||
mask = mask_candidates[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks.append(masks_list)
|
||||
|
||||
# Create 3D mask structure with fixed dimensions (t=50, l=60)
|
||||
masks_3d: List[List[List[List[int]]]] = []
|
||||
for i in range(time_step_num): # Fixed t dimension = 50
|
||||
row = []
|
||||
for j in range(layer_num): # Fixed l dimension = 60
|
||||
row.append(selected_masks) # Add all masks at each position
|
||||
masks_3d.append(row)
|
||||
|
||||
return masks_3d
|
||||
|
||||
elif mode == 'STA_tuning':
|
||||
# Get required parameters
|
||||
mask_search_files_path: Optional[str] = kwargs.get(
|
||||
'mask_search_files_path')
|
||||
if not mask_search_files_path:
|
||||
raise ValueError(
|
||||
"mask_search_files_path is required for STA_tuning mode")
|
||||
|
||||
# Get optional parameters with defaults
|
||||
mask_candidates_tuning: Optional[List[str]] = kwargs.get(
|
||||
'mask_candidates')
|
||||
if mask_candidates_tuning is None:
|
||||
raise ValueError("mask_candidates is required for STA_tuning mode")
|
||||
mask_selected_tuning: List[int] = kwargs.get(
|
||||
'mask_selected', list(range(len(mask_candidates_tuning))))
|
||||
skip_time_steps_tuning: Optional[int] = kwargs.get('skip_time_steps')
|
||||
save_dir_tuning: Optional[str] = kwargs.get('save_dir',
|
||||
"mask_candidates")
|
||||
|
||||
# Parse selected masks
|
||||
selected_masks_tuning: List[List[int]] = []
|
||||
for index in mask_selected_tuning:
|
||||
mask = mask_candidates_tuning[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks_tuning.append(masks_list)
|
||||
|
||||
# Read JSON results
|
||||
results = read_specific_json_files(mask_search_files_path)
|
||||
averaged_results = average_head_losses(results, selected_masks_tuning)
|
||||
|
||||
# Add full attention mask for specific cases
|
||||
full_attention_mask_tuning: Optional[List[int]] = kwargs.get(
|
||||
'full_attention_mask')
|
||||
if full_attention_mask_tuning is not None:
|
||||
selected_masks_tuning.append(full_attention_mask_tuning)
|
||||
|
||||
# Select best mask strategy
|
||||
timesteps_tuning: int = kwargs.get('timesteps', time_step_num)
|
||||
if skip_time_steps_tuning is None:
|
||||
skip_time_steps_tuning = 12
|
||||
mask_strategy, sparsity, strategy_counts = select_best_mask_strategy(
|
||||
averaged_results, selected_masks_tuning, skip_time_steps_tuning,
|
||||
timesteps_tuning, head_num)
|
||||
|
||||
# Save mask strategy
|
||||
if save_dir_tuning is not None:
|
||||
os.makedirs(save_dir_tuning, exist_ok=True)
|
||||
file_path = os.path.join(
|
||||
save_dir_tuning,
|
||||
f'mask_strategy_s{skip_time_steps_tuning}.json')
|
||||
with open(file_path, 'w') as f:
|
||||
json.dump(mask_strategy, f, indent=4)
|
||||
print(f"Successfully saved mask_strategy to {file_path}")
|
||||
|
||||
# Print sparsity and strategy counts for information
|
||||
print(f"Overall sparsity: {sparsity:.4f}")
|
||||
print("\nStrategy usage counts:")
|
||||
total_heads = time_step_num * layer_num * head_num # Fixed dimensions
|
||||
for strategy, count in strategy_counts.items():
|
||||
print(
|
||||
f"Strategy {strategy}: {count} heads ({count/total_heads*100:.2f}%)"
|
||||
)
|
||||
|
||||
# Convert dictionary to 3D list with fixed dimensions
|
||||
mask_strategy_3d = dict_to_3d_list(mask_strategy,
|
||||
t_max=time_step_num,
|
||||
l_max=layer_num,
|
||||
h_max=head_num)
|
||||
|
||||
return mask_strategy_3d
|
||||
elif mode == 'STA_tuning_cfg':
|
||||
# Get required parameters for both positive and negative paths
|
||||
mask_search_files_path_pos: Optional[str] = kwargs.get(
|
||||
'mask_search_files_path_pos')
|
||||
mask_search_files_path_neg: Optional[str] = kwargs.get(
|
||||
'mask_search_files_path_neg')
|
||||
save_dir_cfg: Optional[str] = kwargs.get('save_dir')
|
||||
|
||||
if not mask_search_files_path_pos or not mask_search_files_path_neg or not save_dir_cfg:
|
||||
raise ValueError(
|
||||
"mask_search_files_path_pos, mask_search_files_path_neg, and save_dir are required for STA_tuning_cfg mode"
|
||||
)
|
||||
|
||||
# Get optional parameters with defaults
|
||||
mask_candidates_cfg: Optional[List[str]] = kwargs.get('mask_candidates')
|
||||
if mask_candidates_cfg is None:
|
||||
raise ValueError(
|
||||
"mask_candidates is required for STA_tuning_cfg mode")
|
||||
mask_selected_cfg: List[int] = kwargs.get(
|
||||
'mask_selected', list(range(len(mask_candidates_cfg))))
|
||||
skip_time_steps_cfg: Optional[int] = kwargs.get('skip_time_steps')
|
||||
|
||||
# Parse selected masks
|
||||
selected_masks_cfg: List[List[int]] = []
|
||||
for index in mask_selected_cfg:
|
||||
mask = mask_candidates_cfg[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks_cfg.append(masks_list)
|
||||
|
||||
# Read JSON results for both positive and negative paths
|
||||
pos_results = read_specific_json_files(mask_search_files_path_pos)
|
||||
neg_results = read_specific_json_files(mask_search_files_path_neg)
|
||||
# Combine positive and negative results into one list
|
||||
combined_results = pos_results + neg_results
|
||||
|
||||
# Average the combined results
|
||||
averaged_results = average_head_losses(combined_results,
|
||||
selected_masks_cfg)
|
||||
|
||||
# Add full attention mask for specific cases
|
||||
full_attention_mask_cfg: Optional[List[int]] = kwargs.get(
|
||||
'full_attention_mask')
|
||||
if full_attention_mask_cfg is not None:
|
||||
selected_masks_cfg.append(full_attention_mask_cfg)
|
||||
|
||||
timesteps_cfg: int = kwargs.get('timesteps', time_step_num)
|
||||
if skip_time_steps_cfg is None:
|
||||
skip_time_steps_cfg = 12
|
||||
# Select best mask strategy using combined results
|
||||
mask_strategy, sparsity, strategy_counts = select_best_mask_strategy(
|
||||
averaged_results, selected_masks_cfg, skip_time_steps_cfg,
|
||||
timesteps_cfg, head_num)
|
||||
|
||||
# Save mask strategy
|
||||
os.makedirs(save_dir_cfg, exist_ok=True)
|
||||
file_path = os.path.join(save_dir_cfg,
|
||||
f'mask_strategy_s{skip_time_steps_cfg}.json')
|
||||
with open(file_path, 'w') as f:
|
||||
json.dump(mask_strategy, f, indent=4)
|
||||
print(f"Successfully saved mask_strategy to {file_path}")
|
||||
|
||||
# Print sparsity and strategy counts for information
|
||||
print(f"Overall sparsity: {sparsity:.4f}")
|
||||
print("\nStrategy usage counts:")
|
||||
total_heads = time_step_num * layer_num * head_num # Fixed dimensions
|
||||
for strategy, count in strategy_counts.items():
|
||||
print(
|
||||
f"Strategy {strategy}: {count} heads ({count/total_heads*100:.2f}%)"
|
||||
)
|
||||
|
||||
# Convert dictionary to 3D list with fixed dimensions
|
||||
mask_strategy_3d = dict_to_3d_list(mask_strategy,
|
||||
t_max=time_step_num,
|
||||
l_max=layer_num,
|
||||
h_max=head_num)
|
||||
|
||||
return mask_strategy_3d
|
||||
|
||||
else: # STA_inference
|
||||
# Get parameters with defaults
|
||||
load_path: Optional[str] = kwargs.get(
|
||||
'load_path', "mask_candidates/mask_strategy.json")
|
||||
if load_path is None:
|
||||
raise ValueError("load_path is required for STA_inference mode")
|
||||
|
||||
# Load previously saved mask strategy
|
||||
with open(load_path) as f:
|
||||
mask_strategy = json.load(f)
|
||||
|
||||
# Convert dictionary to 3D list with fixed dimensions
|
||||
mask_strategy_3d = dict_to_3d_list(mask_strategy,
|
||||
t_max=time_step_num,
|
||||
l_max=layer_num,
|
||||
h_max=head_num)
|
||||
|
||||
return mask_strategy_3d
|
||||
|
||||
|
||||
# Helper functions
|
||||
|
||||
|
||||
def read_specific_json_files(folder_path: str) -> List[Dict[str, Any]]:
|
||||
"""Read and parse JSON files containing mask search results."""
|
||||
json_contents: List[Dict[str, Any]] = []
|
||||
|
||||
# List files only in the current directory (no walk)
|
||||
files = os.listdir(folder_path)
|
||||
# Filter files
|
||||
matching_files = [f for f in files if 'mask' in f and f.endswith('.json')]
|
||||
print(f"Found {len(matching_files)} matching files: {matching_files}")
|
||||
|
||||
for file_name in matching_files:
|
||||
file_path = os.path.join(folder_path, file_name)
|
||||
with open(file_path) as file:
|
||||
data = json.load(file)
|
||||
json_contents.append(data)
|
||||
|
||||
return json_contents
|
||||
|
||||
|
||||
def average_head_losses(
|
||||
results: List[Dict[str, Any]],
|
||||
selected_masks: List[List[int]]) -> Dict[str, Dict[str, np.ndarray]]:
|
||||
"""Average losses across all prompts for each mask strategy."""
|
||||
# Initialize a dictionary to store the averaged results
|
||||
averaged_losses: Dict[str, Dict[str, np.ndarray]] = {}
|
||||
loss_type = 'L2_loss'
|
||||
# Get all loss types (e.g., 'L2_loss')
|
||||
averaged_losses[loss_type] = {}
|
||||
|
||||
for mask in selected_masks:
|
||||
mask_str = str(mask)
|
||||
data_shape = np.array(results[0][loss_type][mask_str]).shape
|
||||
accumulated_data = np.zeros(data_shape)
|
||||
|
||||
# Sum across all prompts
|
||||
for prompt_result in results:
|
||||
accumulated_data += np.array(prompt_result[loss_type][mask_str])
|
||||
|
||||
# Average by dividing by number of prompts
|
||||
averaged_data = accumulated_data / len(results)
|
||||
averaged_losses[loss_type][mask_str] = averaged_data
|
||||
|
||||
return averaged_losses
|
||||
|
||||
|
||||
def select_best_mask_strategy(
|
||||
averaged_results: Dict[str, Dict[str, np.ndarray]],
|
||||
selected_masks: List[List[int]],
|
||||
skip_time_steps: int = 12,
|
||||
timesteps: int = 50,
|
||||
head_num: int = 40
|
||||
) -> Tuple[Dict[str, List[int]], float, Dict[str, int]]:
|
||||
"""Select the best mask strategy for each head based on loss minimization."""
|
||||
best_mask_strategy: Dict[str, List[int]] = {}
|
||||
loss_type = 'L2_loss'
|
||||
# Get the shape of time steps and layers
|
||||
layers = len(averaged_results[loss_type][str(selected_masks[0])][0])
|
||||
|
||||
# Counter for sparsity calculation
|
||||
total_tokens = 0 # total number of masked tokens
|
||||
total_length = 0 # total sequence length
|
||||
|
||||
strategy_counts: Dict[str, int] = {
|
||||
str(strategy): 0
|
||||
for strategy in selected_masks
|
||||
}
|
||||
full_attn_strategy = selected_masks[-1] # Last strategy is full attention
|
||||
print(f"Strategy {full_attn_strategy}, skip first {skip_time_steps} steps ")
|
||||
|
||||
for t in range(timesteps):
|
||||
for layer_idx in range(layers):
|
||||
for h in range(head_num):
|
||||
if t < skip_time_steps: # First steps use full attention
|
||||
strategy = full_attn_strategy
|
||||
else:
|
||||
# Get losses for this head across all strategies
|
||||
head_losses = []
|
||||
for strategy in selected_masks[:
|
||||
-1]: # Exclude full attention
|
||||
head_losses.append(averaged_results[loss_type][str(
|
||||
strategy)][t][layer_idx][h])
|
||||
|
||||
# Find which strategy gives minimum loss
|
||||
best_strategy_idx = np.argmin(head_losses)
|
||||
strategy = selected_masks[best_strategy_idx]
|
||||
|
||||
best_mask_strategy[f'{t}_{layer_idx}_{h}'] = strategy
|
||||
|
||||
# Calculate sparsity
|
||||
nums = strategy # strategy is already a list of numbers
|
||||
total_tokens += nums[0] * nums[1] * nums[
|
||||
2] # masked tokens for chosen strategy
|
||||
total_length += full_attn_strategy[0] * full_attn_strategy[
|
||||
1] * full_attn_strategy[2]
|
||||
|
||||
# Count strategy usage
|
||||
strategy_counts[str(strategy)] += 1
|
||||
|
||||
overall_sparsity = 1 - total_tokens / total_length
|
||||
|
||||
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,
|
||||
mask_strategies: List[str],
|
||||
output_dir: str = 'output/mask_search_result/') -> Optional[str]:
|
||||
if not mask_search_final_result:
|
||||
print("No mask search results to save")
|
||||
return None
|
||||
|
||||
# Create result dictionary with defaultdict for nested lists
|
||||
mask_search_dict: Dict[str, Dict[str, List[List[float]]]] = {
|
||||
"L2_loss": defaultdict(list),
|
||||
"L1_loss": defaultdict(list)
|
||||
}
|
||||
|
||||
mask_selected = list(range(len(mask_strategies)))
|
||||
selected_masks: List[List[int]] = []
|
||||
for index in mask_selected:
|
||||
mask = mask_strategies[index]
|
||||
masks_list = [int(x) for x in mask.split(',')]
|
||||
selected_masks.append(masks_list)
|
||||
|
||||
# Process each mask strategy
|
||||
for i, mask_strategy in enumerate(selected_masks):
|
||||
mask_strategy_str = str(mask_strategy)
|
||||
# Process L2 loss
|
||||
step_results: List[List[float]] = []
|
||||
for step_data in mask_search_final_result:
|
||||
if isinstance(step_data, dict) and "L2_loss" in step_data:
|
||||
layer_losses = [float(loss) for loss in step_data["L2_loss"]]
|
||||
step_results.append(layer_losses)
|
||||
mask_search_dict["L2_loss"][mask_strategy_str] = step_results
|
||||
|
||||
step_results = []
|
||||
for step_data in mask_search_final_result:
|
||||
if isinstance(step_data, dict) and "L1_loss" in step_data:
|
||||
layer_losses = [float(loss) for loss in step_data["L1_loss"]]
|
||||
step_results.append(layer_losses)
|
||||
mask_search_dict["L1_loss"][mask_strategy_str] = step_results
|
||||
|
||||
# Create the output directory if it doesn't exist
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
# Create a filename based on the first 20 characters of the prompt
|
||||
filename = prompt[:50].replace(" ", "_")
|
||||
filepath = os.path.join(output_dir, f'mask_search_{filename}.json')
|
||||
|
||||
# Save the results to a JSON file
|
||||
with open(filepath, 'w') as f:
|
||||
json.dump(mask_search_dict, f, indent=4)
|
||||
|
||||
print(f"Successfully saved mask research results to {filepath}")
|
||||
|
||||
return filepath
|
||||
@@ -1,17 +0,0 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
|
||||
AttentionMetadata,
|
||||
AttentionMetadataBuilder)
|
||||
from fastvideo.v1.attention.layer import DistributedAttention, LocalAttention
|
||||
from fastvideo.v1.attention.selector import get_attn_backend
|
||||
|
||||
__all__ = [
|
||||
"DistributedAttention",
|
||||
"LocalAttention",
|
||||
"AttentionBackend",
|
||||
"AttentionMetadata",
|
||||
"AttentionMetadataBuilder",
|
||||
# "AttentionState",
|
||||
"get_attn_backend",
|
||||
]
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field, fields
|
||||
from typing import Any, Dict
|
||||
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Optional, Tuple
|
||||
from typing import Any, List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
|
||||
from fastvideo.v1.layers.quantization import QuantizationConfig
|
||||
@@ -11,15 +12,18 @@ class DiTArchConfig(ArchConfig):
|
||||
_fsdp_shard_conditions: list = field(default_factory=list)
|
||||
_compile_conditions: list = field(default_factory=list)
|
||||
_param_names_mapping: dict = field(default_factory=dict)
|
||||
_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.TORCH_SDPA,
|
||||
_Backend.VIDEO_SPARSE_ATTN)
|
||||
|
||||
hidden_size: int = 0
|
||||
num_attention_heads: int = 0
|
||||
num_channels_latents: int = 0
|
||||
exclude_lora_layers: List[str] = field(default_factory=list)
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self._compile_conditions:
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Tuple
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
import torch
|
||||
|
||||
@@ -163,6 +164,8 @@ class HunyuanVideoArchConfig(DiTArchConfig):
|
||||
pooled_projection_dim: int = 768
|
||||
rope_theta: int = 256
|
||||
qk_norm: str = "rms_norm"
|
||||
exclude_lora_layers: List[str] = field(
|
||||
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Optional, Tuple, Union
|
||||
|
||||
@@ -51,6 +52,7 @@ class StepVideoArchConfig(DiTArchConfig):
|
||||
default_factory=lambda: [6144, 1024])
|
||||
attention_type: Optional[str] = "torch"
|
||||
use_additional_conditions: Optional[bool] = False
|
||||
exclude_lora_layers: List[str] = field(default_factory=lambda: [])
|
||||
|
||||
def __post_init__(self):
|
||||
self.hidden_size = self.num_attention_heads * self.attention_head_dim
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Tuple
|
||||
from typing import List, Optional, Tuple
|
||||
|
||||
from fastvideo.v1.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
@@ -51,6 +52,23 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
r"blocks\.(\d+)\.norm2\.(.*)$":
|
||||
r"blocks.\1.self_attn_residual_norm.norm.\2",
|
||||
})
|
||||
# Some LoRA adapters use the original official layer names instead of hf layer names,
|
||||
# so apply this before the param_names_mapping
|
||||
_lora_param_names_mapping: dict = field(
|
||||
default_factory=lambda: {
|
||||
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.attn1.to_q.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.attn1.to_k.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.attn1.to_v.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$":
|
||||
r"blocks.\1.attn1.to_out.0.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$":
|
||||
r"blocks.\1.attn2.to_out.0.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
|
||||
})
|
||||
|
||||
patch_size: Tuple[int, int, int] = (1, 2, 2)
|
||||
text_len = 512
|
||||
@@ -68,6 +86,7 @@ class WanVideoArchConfig(DiTArchConfig):
|
||||
image_dim: Optional[int] = None
|
||||
added_kv_proj_dim: Optional[int] = None
|
||||
rope_max_seq_len: int = 1024
|
||||
exclude_lora_layers: List[str] = field(default_factory=lambda: ["embedder"])
|
||||
|
||||
def __post_init__(self):
|
||||
super().__post_init__()
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional
|
||||
|
||||
|
||||
@@ -1,3 +1,6 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import argparse
|
||||
import dataclasses
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Union
|
||||
|
||||
@@ -128,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)
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Tuple
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Tuple
|
||||
|
||||
|
||||
@@ -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,14 +1,17 @@
|
||||
# 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 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__)
|
||||
|
||||
@@ -21,58 +24,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
|
||||
use_cpu_offload: bool = False
|
||||
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: Optional[str] = None
|
||||
skip_time_steps: int = 15
|
||||
|
||||
# Compilation
|
||||
enable_torch_compile: bool = False
|
||||
# enable_torch_compile: bool = False
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser,
|
||||
prefix: str = "") -> FlexibleArgumentParser:
|
||||
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
|
||||
|
||||
# model_path will be conflicting with the model_path in FastVideoArgs,
|
||||
# so we add it separately if prefix is not empty
|
||||
if prefix_with_dot != "":
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}model-path",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}model_path",
|
||||
default=PipelineConfig.model_path,
|
||||
help="Path to the pretrained model",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}pipeline-config-path",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}pipeline_config_path",
|
||||
default=PipelineConfig.pipeline_config_path,
|
||||
help="Path to the pipeline config",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}embedded-cfg-scale",
|
||||
type=float,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}embedded_cfg_scale",
|
||||
default=PipelineConfig.embedded_cfg_scale,
|
||||
help="Embedded CFG scale",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}flow-shift",
|
||||
type=float,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}flow_shift",
|
||||
default=PipelineConfig.flow_shift,
|
||||
help="Flow shift parameter",
|
||||
)
|
||||
|
||||
# DiT configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}dit-precision",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}dit_precision",
|
||||
default=PipelineConfig.dit_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for the DiT model",
|
||||
)
|
||||
|
||||
# VAE configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}vae-precision",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}vae_precision",
|
||||
default=PipelineConfig.vae_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for VAE",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}vae-tiling",
|
||||
action=StoreBoolean,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}vae_tiling",
|
||||
default=PipelineConfig.vae_tiling,
|
||||
help="Enable VAE tiling",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}vae-sp",
|
||||
action=StoreBoolean,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}vae_sp",
|
||||
help="Enable VAE spatial parallelism",
|
||||
)
|
||||
|
||||
# Text encoder configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}text-encoder-precisions",
|
||||
nargs="+",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}text_encoder_precisions",
|
||||
default=PipelineConfig.DEFAULT_TEXT_ENCODER_PRECISIONS,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for each text encoder",
|
||||
)
|
||||
|
||||
# Image encoder configuration
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}image-encoder-precision",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}image_encoder_precision",
|
||||
default=PipelineConfig.image_encoder_precision,
|
||||
choices=["fp32", "fp16", "bf16"],
|
||||
help="Precision for image encoder",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}pos_magic",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}pos_magic",
|
||||
default=PipelineConfig.pos_magic,
|
||||
help="Positive magic prompt for sampling, used in stepvideo",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}neg_magic",
|
||||
type=str,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}neg_magic",
|
||||
default=PipelineConfig.neg_magic,
|
||||
help="Negative magic prompt for sampling, used in stepvideo",
|
||||
)
|
||||
parser.add_argument(
|
||||
f"--{prefix_with_dot}timesteps_scale",
|
||||
type=bool,
|
||||
dest=f"{prefix_with_dot.replace('-', '_')}timesteps_scale",
|
||||
default=PipelineConfig.timesteps_scale,
|
||||
help=
|
||||
"Bool for applying scheduler scale in set_timesteps, used in stepvideo",
|
||||
)
|
||||
|
||||
# Add VAE configuration arguments
|
||||
from fastvideo.v1.configs.models.vaes.base import VAEConfig
|
||||
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
|
||||
|
||||
# Add DiT configuration arguments
|
||||
from fastvideo.v1.configs.models.dits.base import DiTConfig
|
||||
DiTConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}dit-config")
|
||||
|
||||
return parser
|
||||
|
||||
def update_config_from_dict(self,
|
||||
args: Dict[str, Any],
|
||||
prefix: str = "") -> None:
|
||||
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
|
||||
update_config_from_args(self, args, prefix, pop_args=True)
|
||||
update_config_from_args(self.vae_config,
|
||||
args,
|
||||
f"{prefix_with_dot}vae_config",
|
||||
pop_args=True)
|
||||
update_config_from_args(self.dit_config,
|
||||
args,
|
||||
f"{prefix_with_dot}dit_config",
|
||||
pop_args=True)
|
||||
|
||||
@classmethod
|
||||
def from_pretrained(cls, model_path: str) -> "PipelineConfig":
|
||||
"""
|
||||
use the pipeline class setting from model_path to match the pipeline config
|
||||
"""
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_for_name)
|
||||
pipeline_config_cls = get_pipeline_config_cls_for_name(model_path)
|
||||
if pipeline_config_cls is not None:
|
||||
pipeline_config = pipeline_config_cls()
|
||||
else:
|
||||
get_pipeline_config_cls_from_name)
|
||||
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
|
||||
|
||||
return cast(PipelineConfig, pipeline_config_cls(model_path=model_path))
|
||||
|
||||
@classmethod
|
||||
def from_kwargs(cls,
|
||||
kwargs: Dict[str, Any],
|
||||
config_cli_prefix: str = "") -> "PipelineConfig":
|
||||
"""
|
||||
Load PipelineConfig from kwargs Dictionary.
|
||||
kwargs: dictionary of kwargs
|
||||
config_cli_prefix: prefix of CLI arguments for this PipelineConfig instance
|
||||
"""
|
||||
from fastvideo.v1.configs.pipelines.registry import (
|
||||
get_pipeline_config_cls_from_name)
|
||||
|
||||
prefix_with_dot = f"{config_cli_prefix}." if (config_cli_prefix.strip()
|
||||
!= "") else ""
|
||||
model_path: Optional[str] = kwargs.get(prefix_with_dot + 'model_path',
|
||||
None) or kwargs.get('model_path')
|
||||
pipeline_config_or_path: Optional[Union[str, PipelineConfig, Dict[
|
||||
str, Any]]] = kwargs.get(prefix_with_dot + 'pipeline_config',
|
||||
None) or kwargs.get('pipeline_config')
|
||||
if model_path is None:
|
||||
raise ValueError("model_path is required in kwargs")
|
||||
|
||||
# 1. Get the pipeline config class from the registry
|
||||
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
|
||||
|
||||
# 2. Instantiate PipelineConfig
|
||||
if pipeline_config_cls is None:
|
||||
logger.warning(
|
||||
"Couldn't find an optimal sampling param for %s. Using the default sampling param.",
|
||||
"Couldn't find pipeline config for %s. Using the default pipeline config.",
|
||||
model_path)
|
||||
pipeline_config = cls()
|
||||
else:
|
||||
pipeline_config = pipeline_config_cls()
|
||||
|
||||
return cast(PipelineConfig, pipeline_config)
|
||||
# 3. Load PipelineConfig from a json file or a PipelineConfig object if provided
|
||||
if isinstance(pipeline_config_or_path, str):
|
||||
pipeline_config.load_from_json(pipeline_config_or_path)
|
||||
kwargs[prefix_with_dot +
|
||||
'pipeline_config_path'] = pipeline_config_or_path
|
||||
elif isinstance(pipeline_config_or_path, PipelineConfig):
|
||||
pipeline_config = pipeline_config_or_path
|
||||
elif isinstance(pipeline_config_or_path, dict):
|
||||
pipeline_config.update_pipeline_config(pipeline_config_or_path)
|
||||
|
||||
# 4. Update PipelineConfig from CLI arguments if provided
|
||||
kwargs[prefix_with_dot + 'model_path'] = model_path
|
||||
pipeline_config.update_config_from_dict(kwargs, config_cli_prefix)
|
||||
return pipeline_config
|
||||
|
||||
def check_pipeline_config(self) -> None:
|
||||
if self.vae_sp and not self.vae_tiling:
|
||||
raise ValueError(
|
||||
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
|
||||
)
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.text_encoder_precisions):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
|
||||
)
|
||||
|
||||
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
)
|
||||
|
||||
if len(self.preprocess_text_funcs) != len(self.postprocess_text_funcs):
|
||||
raise ValueError(
|
||||
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
|
||||
)
|
||||
|
||||
def dump_to_json(self, file_path: str):
|
||||
output_dict = shallow_asdict(self)
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Tuple, TypedDict
|
||||
|
||||
@@ -68,9 +69,6 @@ class HunyuanConfig(PipelineConfig):
|
||||
embedded_cfg_scale: int = 6
|
||||
flow_shift: int = 7
|
||||
|
||||
# Video parameters
|
||||
use_cpu_offload: bool = True
|
||||
|
||||
# Text encoding stage
|
||||
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda: (LlamaConfig(), CLIPTextConfig()))
|
||||
@@ -82,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"))
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Registry for pipeline weight-specific configurations."""
|
||||
|
||||
import os
|
||||
@@ -18,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,
|
||||
@@ -50,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
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.models import DiTConfig, VAEConfig
|
||||
@@ -18,9 +19,6 @@ class StepVideoT2VConfig(PipelineConfig):
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
# Video parameters
|
||||
use_cpu_offload: bool = True
|
||||
|
||||
# Denoising stage
|
||||
flow_shift: int = 13
|
||||
timesteps_scale: bool = False
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Tuple
|
||||
|
||||
@@ -37,9 +38,6 @@ class WanT2V480PConfig(PipelineConfig):
|
||||
vae_tiling: bool = False
|
||||
vae_sp: bool = False
|
||||
|
||||
# Video parameters
|
||||
use_cpu_offload: bool = True
|
||||
|
||||
# Denoising stage
|
||||
flow_shift: int = 3
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict, List, Optional, Union
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
from typing import Any, Callable, Dict, Optional
|
||||
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.sample.base import CacheParams
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.v1.configs.sample.base import SamplingParam
|
||||
|
||||
@@ -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)
|
||||
@@ -8,6 +8,10 @@ from fastvideo.v1.dataset.t2v_datasets import T2V_dataset
|
||||
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"]
|
||||
|
||||
|
||||
def getdataset(args, start_idx=0) -> T2V_dataset:
|
||||
temporal_sample = TemporalRandomCrop(args.num_frames) # 16 x
|
||||
|
||||
@@ -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()
|
||||
@@ -0,0 +1,187 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import argparse
|
||||
import os
|
||||
import pathlib
|
||||
import time
|
||||
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.distributed.checkpoint as dist_cp
|
||||
|
||||
from fastvideo.v1.dataset.parquet_dataset_map_style import (
|
||||
build_parquet_map_style_dataloader)
|
||||
from fastvideo.v1.distributed import get_world_rank
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_torch_device,
|
||||
maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
torch.multiprocessing.set_start_method("spawn", force=True)
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Benchmark parquet map 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_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))
|
||||
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_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,
|
||||
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_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,
|
||||
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
|
||||
|
||||
dataset, dataloader = build_parquet_map_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) 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()
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# schema.py
|
||||
"""
|
||||
Unified data schema and format for saving and loading image/video data after
|
||||
@@ -9,7 +10,45 @@ frameworks that can handle parquet or lance file.
|
||||
|
||||
import pyarrow as pa
|
||||
|
||||
pyarrow_schema = pa.schema([
|
||||
pyarrow_schema_i2v = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Image/Video VAE latents ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("vae_latent_bytes", pa.binary()),
|
||||
# e.g., [C, T, H, W] or [C, H, W]
|
||||
pa.field("vae_latent_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'float32'
|
||||
pa.field("vae_latent_dtype", pa.string()),
|
||||
# --- Text encoder output tensor ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
pa.field("text_embedding_bytes", pa.binary()),
|
||||
# e.g., [SeqLen, Dim]
|
||||
pa.field("text_embedding_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bfloat16' or 'float32'
|
||||
pa.field("text_embedding_dtype", pa.string()),
|
||||
pa.field("text_attention_mask_bytes", pa.binary()),
|
||||
# e.g., [SeqLen]
|
||||
pa.field("text_attention_mask_shape", pa.list_(pa.int64())),
|
||||
# e.g., 'bool' or 'int8'
|
||||
pa.field("text_attention_mask_dtype", pa.string()),
|
||||
#I2V
|
||||
pa.field("clip_feature_bytes", pa.binary()),
|
||||
pa.field("clip_feature_shape", pa.list_(pa.int64())),
|
||||
pa.field("clip_feature_dtype", pa.string()),
|
||||
# --- Metadata ---
|
||||
pa.field("file_name", pa.string()),
|
||||
pa.field("caption", pa.string()),
|
||||
pa.field("media_type", pa.string()), # 'image' or 'video'
|
||||
pa.field("width", pa.int64()),
|
||||
pa.field("height", pa.int64()),
|
||||
# -- Video-specific (can be null/default for images) ---
|
||||
# Number of frames processed (e.g., 1 for image, N for video)
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
|
||||
pyarrow_schema_t2v = pa.schema([
|
||||
pa.field("id", pa.string()),
|
||||
# --- Image/Video VAE latents ---
|
||||
# Tensors are stored as raw bytes with shape and dtype info for loading
|
||||
@@ -41,4 +80,4 @@ pyarrow_schema = pa.schema([
|
||||
pa.field("num_frames", pa.int64()),
|
||||
pa.field("duration_sec", pa.float64()),
|
||||
pa.field("fps", pa.float64()),
|
||||
])
|
||||
])
|
||||
@@ -0,0 +1,137 @@
|
||||
# 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)
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
@@ -107,23 +108,3 @@ def latent_collate_function(batch):
|
||||
prompt_attention_masks = torch.stack(prompt_attention_masks, dim=0)
|
||||
latents = torch.stack(latent_list, dim=0)
|
||||
return latents, prompt_embeds, latent_attn_mask, prompt_attention_masks
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
dataset = LatentDataset("data/Mochi-Synthetic-Data/merge.txt",
|
||||
num_latent_t=28,
|
||||
cfg_rate=0.0)
|
||||
dataloader = torch.utils.data.DataLoader(dataset,
|
||||
batch_size=2,
|
||||
shuffle=False,
|
||||
collate_fn=latent_collate_function)
|
||||
for latent, prompt_embed, latent_attn_mask, prompt_attention_mask in dataloader:
|
||||
print(
|
||||
latent.shape,
|
||||
prompt_embed.shape,
|
||||
latent_attn_mask.shape,
|
||||
prompt_attention_mask.shape,
|
||||
)
|
||||
import pdb
|
||||
|
||||
pdb.set_trace()
|
||||
|
||||
@@ -0,0 +1,275 @@
|
||||
import os
|
||||
import pickle
|
||||
import random
|
||||
from typing import Dict, List, Tuple
|
||||
|
||||
import numpy as np
|
||||
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):
|
||||
super().__init__()
|
||||
self.path = str(path)
|
||||
self.batch_size = batch_size
|
||||
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
|
||||
@@ -0,0 +1,301 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
import pickle
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
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_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 DP_SP_BatchSampler(Sampler[List[int]]):
|
||||
"""
|
||||
A simple sequential batch sampler that yields batches of indices.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
batch_size: int,
|
||||
dataset_size: int,
|
||||
num_sp_groups: int,
|
||||
sp_world_size: int,
|
||||
global_rank: int,
|
||||
drop_last: bool = True,
|
||||
seed: int = 0,
|
||||
):
|
||||
self.batch_size = batch_size
|
||||
self.dataset_size = dataset_size
|
||||
self.drop_last = drop_last
|
||||
self.seed = seed
|
||||
self.num_sp_groups = num_sp_groups
|
||||
self.global_rank = global_rank
|
||||
self.sp_world_size = sp_world_size
|
||||
|
||||
# ── epoch-level RNG ────────────────────────────────────────────────
|
||||
rng = torch.Generator().manual_seed(self.seed)
|
||||
# Create a random permutation of all indices
|
||||
global_indices = torch.randperm(self.dataset_size, generator=rng)
|
||||
|
||||
if self.drop_last:
|
||||
# For drop_last=True, we:
|
||||
# 1. Ensure total samples is divisible by (batch_size * num_sp_groups)
|
||||
# 2. This guarantees each SP group gets same number of complete batches
|
||||
# 3. Prevents uneven batch sizes across SP groups at end of epoch
|
||||
num_batches = self.dataset_size // self.batch_size
|
||||
num_global_batches = num_batches // self.num_sp_groups
|
||||
global_indices = global_indices[:num_global_batches *
|
||||
self.num_sp_groups *
|
||||
self.batch_size]
|
||||
else:
|
||||
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("Dataset size for each sp group: %d",
|
||||
len(sp_group_local_indices))
|
||||
|
||||
def __iter__(self):
|
||||
indices = self.sp_group_local_indices
|
||||
for i in range(0, len(indices), self.batch_size):
|
||||
batch_indices = indices[i:i + self.batch_size]
|
||||
yield batch_indices.tolist()
|
||||
|
||||
def __len__(self):
|
||||
return len(self.sp_group_local_indices) // self.batch_size
|
||||
|
||||
|
||||
def get_parquet_files_and_length(path: str):
|
||||
# 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,
|
||||
lengths: List[int]) -> Dict[str, Any]:
|
||||
'''
|
||||
Read a row from a parquet file.
|
||||
Args:
|
||||
parquet_files: List[str]
|
||||
global_row_idx: int
|
||||
lengths: List[int]
|
||||
Returns:
|
||||
'''
|
||||
# find the parquet file and local row index
|
||||
cumulative = 0
|
||||
for file_index in range(len(lengths)):
|
||||
if cumulative + lengths[file_index] > global_row_idx:
|
||||
local_row_idx = global_row_idx - cumulative
|
||||
break
|
||||
cumulative += lengths[file_index]
|
||||
|
||||
parquet_file = pq.ParquetFile(parquet_files[file_index])
|
||||
|
||||
# Calculate the row group to read into memory and the local idx
|
||||
# This way we can avoid reading in the entire parquet file
|
||||
cumulative = 0
|
||||
for i in range(parquet_file.num_row_groups):
|
||||
num_rows = parquet_file.metadata.row_group(i).num_rows
|
||||
if cumulative + num_rows > local_row_idx:
|
||||
row_group_index = i
|
||||
local_index = local_row_idx - cumulative
|
||||
break
|
||||
cumulative += num_rows
|
||||
|
||||
row_group = parquet_file.read_row_group(row_group_index).to_pydict()
|
||||
row_dict = {k: v[local_index] for k, v in row_group.items()}
|
||||
del row_group
|
||||
|
||||
return row_dict
|
||||
|
||||
|
||||
# ────────────────────────────────────────────────────────────────────────────
|
||||
# 2. Dataset with batched __getitems__
|
||||
# ────────────────────────────────────────────────────────────────────────────
|
||||
class LatentsParquetMapStyleDataset(Dataset):
|
||||
"""
|
||||
Return latents[B,C,T,H,W] and embeddings[B,L,D] in pinned CPU memory.
|
||||
Note:
|
||||
Using parquet for map style dataset is not efficient, we mainly keep it for backward compatibility and debugging.
|
||||
"""
|
||||
# Modify this in the future if we want to add more keys, for example, in image to video.
|
||||
keys = [("vae_latent", "latent"), "text_embedding"]
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
path: str,
|
||||
batch_size: int,
|
||||
cfg_rate: float = 0.0,
|
||||
seed: int = 42,
|
||||
drop_last: bool = True,
|
||||
text_padding_length: int = 512,
|
||||
):
|
||||
super().__init__()
|
||||
self.path = path
|
||||
self.cfg_rate = cfg_rate
|
||||
if cfg_rate > 0.0:
|
||||
raise ValueError(
|
||||
"cfg_rate > 0.0 is not supported for now because it will trigger bug when num_data_workers > 0"
|
||||
)
|
||||
logger.info("Initializing LatentsParquetMapStyleDataset with path: %s",
|
||||
path)
|
||||
self.parquet_files, self.lengths = get_parquet_files_and_length(path)
|
||||
self.batch = batch_size
|
||||
self.text_padding_length = text_padding_length
|
||||
self._cols = [
|
||||
"vae_latent_bytes",
|
||||
"vae_latent_shape",
|
||||
"text_embedding_bytes",
|
||||
"text_embedding_shape",
|
||||
"text_embedding_dtype",
|
||||
"height",
|
||||
"width",
|
||||
]
|
||||
self.sampler = DP_SP_BatchSampler(
|
||||
batch_size=batch_size,
|
||||
dataset_size=sum(self.lengths),
|
||||
num_sp_groups=get_world_size() // get_sp_world_size(),
|
||||
sp_world_size=get_sp_world_size(),
|
||||
global_rank=get_world_rank(),
|
||||
drop_last=drop_last,
|
||||
seed=seed,
|
||||
)
|
||||
logger.info("Dataset initialized with %d parquet files and %d rows",
|
||||
len(self.parquet_files), sum(self.lengths))
|
||||
|
||||
def get_validation_negative_prompt(
|
||||
self) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, str]:
|
||||
"""
|
||||
Get the negative prompt for validation.
|
||||
This method ensures the negative prompt is loaded and cached properly.
|
||||
Returns the processed negative prompt data (latents, embeddings, masks, info).
|
||||
"""
|
||||
|
||||
# Read first row from first parquet file
|
||||
file_path = self.parquet_files[0]
|
||||
row_idx = 0
|
||||
# Read the negative prompt data
|
||||
row_dict = read_row_from_parquet_file([file_path], row_idx,
|
||||
[self.lengths[0]])
|
||||
|
||||
all_latents_list, all_embs_list, all_masks_list, caption_text_list = collate_latents_embs_masks(
|
||||
[row_dict], self.text_padding_length, self.keys)
|
||||
all_latents, all_embs, all_masks, caption_text = all_latents_list[
|
||||
0], all_embs_list[0], all_masks_list[0], caption_text_list[0]
|
||||
# add batch dimension
|
||||
if len(all_embs.shape) == 2:
|
||||
all_embs = all_embs.unsqueeze(0)
|
||||
if len(all_masks.shape) == 1:
|
||||
all_masks = all_masks.unsqueeze(0).unsqueeze(0)
|
||||
return all_latents, all_embs, all_masks, caption_text
|
||||
|
||||
# PyTorch calls this ONLY because the batch_sampler yields a list
|
||||
def __getitems__(self, indices: List[int]):
|
||||
"""
|
||||
Batch fetch using read_row_from_parquet_file for each index.
|
||||
"""
|
||||
rows = [
|
||||
read_row_from_parquet_file(self.parquet_files, idx, self.lengths)
|
||||
for idx in indices
|
||||
]
|
||||
|
||||
all_latents, all_embs, all_masks, caption_text = collate_latents_embs_masks(
|
||||
rows, self.text_padding_length, self.keys)
|
||||
return all_latents, all_embs, all_masks, caption_text
|
||||
|
||||
def __len__(self):
|
||||
return sum(self.lengths)
|
||||
|
||||
|
||||
# ────────────────────────────────────────────────────────────────────────────
|
||||
# 3. Loader helper – everything else stays just like your original trainer
|
||||
# ────────────────────────────────────────────────────────────────────────────
|
||||
def passthrough(batch):
|
||||
return batch
|
||||
|
||||
|
||||
def build_parquet_map_style_dataloader(
|
||||
path,
|
||||
batch_size,
|
||||
num_data_workers,
|
||||
cfg_rate=0.0,
|
||||
drop_last=True,
|
||||
text_padding_length=512,
|
||||
seed=42) -> Tuple[LatentsParquetMapStyleDataset, StatefulDataLoader]:
|
||||
dataset = LatentsParquetMapStyleDataset(
|
||||
path,
|
||||
batch_size,
|
||||
cfg_rate=cfg_rate,
|
||||
drop_last=drop_last,
|
||||
text_padding_length=text_padding_length,
|
||||
seed=seed)
|
||||
|
||||
loader = StatefulDataLoader(
|
||||
dataset,
|
||||
batch_sampler=dataset.sampler,
|
||||
collate_fn=passthrough,
|
||||
num_workers=num_data_workers,
|
||||
pin_memory=True,
|
||||
persistent_workers=num_data_workers > 0,
|
||||
)
|
||||
return dataset, loader
|
||||
@@ -1,369 +0,0 @@
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
import time
|
||||
from collections import defaultdict
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import numpy as np
|
||||
import pyarrow.parquet as pq
|
||||
import torch
|
||||
import tqdm
|
||||
from einops import rearrange
|
||||
from torch import distributed as dist
|
||||
from torch.utils.data import Dataset
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
|
||||
get_sp_group)
|
||||
from fastvideo.v1.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class ParquetVideoTextDataset(Dataset):
|
||||
"""Efficient loader for video-text data from a directory of Parquet files."""
|
||||
|
||||
def __init__(self,
|
||||
path: str,
|
||||
batch_size: int = 1024,
|
||||
rank: int = 0,
|
||||
world_size: int = 1,
|
||||
cfg_rate: float = 0.0,
|
||||
num_latent_t: int = 2,
|
||||
seed: int = 0):
|
||||
super().__init__()
|
||||
self.path = str(path)
|
||||
self.batch_size = batch_size
|
||||
self.rank = rank
|
||||
self.local_rank = get_sequence_model_parallel_rank()
|
||||
self.sp_world_size = world_size
|
||||
self.world_size = int(os.getenv("WORLD_SIZE", 1))
|
||||
self.cfg_rate = cfg_rate
|
||||
self.num_latent_t = num_latent_t
|
||||
self.local_indices = None
|
||||
self.plan_output_dir = os.path.join(self.path, "data_plan.json")
|
||||
|
||||
ranks = get_sp_group().ranks
|
||||
group_ranks: List[List] = [[] for _ in range(self.world_size)]
|
||||
torch.distributed.all_gather_object(group_ranks, ranks)
|
||||
|
||||
if rank == 0:
|
||||
# If a plan already exists, then skip creating a new plan
|
||||
# This will be useful when resume training
|
||||
if os.path.exists(self.plan_output_dir):
|
||||
print(f"Using existing plan from {self.plan_output_dir}")
|
||||
return
|
||||
|
||||
# Find all parquet files recursively, and record num_rows for each file
|
||||
print(f"Scanning for parquet files in {self.path}")
|
||||
metadatas = []
|
||||
for root, _, files in os.walk(self.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
|
||||
for row_idx in range(num_rows):
|
||||
metadatas.append((file_path, row_idx))
|
||||
|
||||
# Generate the plan that distribute rows among workers
|
||||
random.seed(seed)
|
||||
random.shuffle(metadatas)
|
||||
|
||||
# Get all sp groups
|
||||
# e.g. if num_gpus = 4, sp_size = 2
|
||||
# group_ranks = [(0, 1), (2, 3)]
|
||||
# We will assign the same batches of data to ranks in the same sp group, and we'll assign different batches to ranks in different sp groups
|
||||
# e.g. plan = {0: [row 1, row 4], 1: [row 1, row 4], 2: [row 2, row 3], 3: [row 2, row 3]}
|
||||
group_ranks_list: List[Any] = list(
|
||||
set(tuple(r) for r in group_ranks))
|
||||
num_sp_groups = len(group_ranks_list)
|
||||
plan = defaultdict(list)
|
||||
for idx, metadata in enumerate(metadatas):
|
||||
sp_group_idx = idx % num_sp_groups
|
||||
for global_rank in group_ranks_list[sp_group_idx]:
|
||||
plan[global_rank].append(metadata)
|
||||
|
||||
with open(self.plan_output_dir, "w") as f:
|
||||
json.dump(plan, f)
|
||||
|
||||
def __len__(self):
|
||||
if self.local_indices is None:
|
||||
try:
|
||||
with open(self.plan_output_dir) as f:
|
||||
plan = json.load(f)
|
||||
self.local_indices = plan[str(self.rank)]
|
||||
except Exception as err:
|
||||
raise Exception(
|
||||
"The data plan hasn't been created yet") from err
|
||||
assert self.local_indices is not None
|
||||
return len(self.local_indices)
|
||||
|
||||
def __getitem__(self, idx):
|
||||
if self.local_indices is None:
|
||||
try:
|
||||
with open(self.plan_output_dir) as f:
|
||||
plan = json.load(f)
|
||||
self.local_indices = plan[self.rank]
|
||||
except Exception as err:
|
||||
raise Exception(
|
||||
"The data plan hasn't been created yet") from err
|
||||
assert self.local_indices is not None
|
||||
file_path, row_idx = self.local_indices[idx]
|
||||
parquet_file = pq.ParquetFile(file_path)
|
||||
|
||||
# Calculate the row group to read into memory and the local idx
|
||||
# This way we can avoid reading in the entire parquet file
|
||||
cumulative = 0
|
||||
for i in range(parquet_file.num_row_groups):
|
||||
num_rows = parquet_file.metadata.row_group(i).num_rows
|
||||
if cumulative + num_rows > idx:
|
||||
row_group_index = i
|
||||
local_index = idx - cumulative
|
||||
break
|
||||
cumulative += num_rows
|
||||
|
||||
row_group = parquet_file.read_row_group(row_group_index).to_pydict()
|
||||
row_dict = {k: v[local_index] for k, v in row_group.items()}
|
||||
del row_group
|
||||
|
||||
processed = self._process_row(row_dict)
|
||||
lat, emb, mask, info = processed["latents"], processed[
|
||||
"embeddings"], processed["masks"], processed["info"]
|
||||
if lat.numel() == 0: # Validation parquet
|
||||
return lat, emb, mask, info
|
||||
else:
|
||||
lat = lat[:, -self.num_latent_t:]
|
||||
if self.sp_world_size > 1:
|
||||
lat = rearrange(lat,
|
||||
"t (n s) h w -> t n s h w",
|
||||
n=self.sp_world_size).contiguous()
|
||||
lat = lat[:, self.local_rank, :, :, :]
|
||||
return lat, emb, mask, info
|
||||
|
||||
def _process_row(self, row) -> Dict[str, Any]:
|
||||
"""Process a PyArrow batch into tensors."""
|
||||
|
||||
vae_latent_bytes = row["vae_latent_bytes"]
|
||||
vae_latent_shape = row["vae_latent_shape"]
|
||||
text_embedding_bytes = row["text_embedding_bytes"]
|
||||
text_embedding_shape = row["text_embedding_shape"]
|
||||
text_attention_mask_bytes = row["text_attention_mask_bytes"]
|
||||
text_attention_mask_shape = row["text_attention_mask_shape"]
|
||||
|
||||
# Process latent
|
||||
if not vae_latent_shape: # No VAE latent is stored. Split is validation
|
||||
lat = np.array([])
|
||||
else:
|
||||
lat = np.frombuffer(vae_latent_bytes,
|
||||
dtype=np.float32).reshape(vae_latent_shape)
|
||||
# Make array writable
|
||||
lat = np.copy(lat)
|
||||
|
||||
if random.random() < self.cfg_rate:
|
||||
emb = np.zeros((512, 4096), dtype=np.float32)
|
||||
else:
|
||||
emb = np.frombuffer(text_embedding_bytes,
|
||||
dtype=np.float32).reshape(text_embedding_shape)
|
||||
# Make array writable
|
||||
emb = np.copy(emb)
|
||||
if emb.shape[0] < 512:
|
||||
padded_emb = np.zeros((512, emb.shape[1]), dtype=np.float32)
|
||||
padded_emb[:emb.shape[0], :] = emb
|
||||
emb = padded_emb
|
||||
elif emb.shape[0] > 512:
|
||||
emb = emb[:512, :]
|
||||
|
||||
# Process mask
|
||||
if len(text_attention_mask_bytes) > 0 and len(
|
||||
text_attention_mask_shape) > 0:
|
||||
msk = np.frombuffer(text_attention_mask_bytes,
|
||||
dtype=np.uint8).astype(np.bool_)
|
||||
msk = msk.reshape(1, -1)
|
||||
# Make array writable
|
||||
msk = np.copy(msk)
|
||||
if msk.shape[1] < 512:
|
||||
padded_msk = np.zeros((1, 512), dtype=np.bool_)
|
||||
padded_msk[:, :msk.shape[1]] = msk
|
||||
msk = padded_msk
|
||||
elif msk.shape[1] > 512:
|
||||
msk = msk[:, :512]
|
||||
else:
|
||||
msk = np.ones((1, 512), dtype=np.bool_)
|
||||
|
||||
# Collect metadata
|
||||
info = {
|
||||
"width": row["width"],
|
||||
"height": row["height"],
|
||||
"num_frames": row["num_frames"],
|
||||
"duration_sec": row["duration_sec"],
|
||||
"fps": row["fps"],
|
||||
"file_name": row["file_name"],
|
||||
"caption": row["caption"],
|
||||
}
|
||||
|
||||
return {
|
||||
"latents": torch.from_numpy(lat),
|
||||
"embeddings": torch.from_numpy(emb),
|
||||
"masks": torch.from_numpy(msk),
|
||||
"info": info
|
||||
}
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = argparse.ArgumentParser(
|
||||
description='Benchmark Parquet dataset loading speed')
|
||||
parser.add_argument('--path',
|
||||
type=str,
|
||||
default="your/dataset/path",
|
||||
help='Path to Parquet dataset')
|
||||
parser.add_argument('--batch_size',
|
||||
type=int,
|
||||
default=4,
|
||||
help='Batch size for DataLoader')
|
||||
parser.add_argument('--num_batches',
|
||||
type=int,
|
||||
default=100,
|
||||
help='Number of batches to benchmark')
|
||||
parser.add_argument('--vae_debug', action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
# Initialize distributed training
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
|
||||
# Initialize CUDA device first
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.set_device(local_rank)
|
||||
device = torch.device(f"cuda:{local_rank}")
|
||||
else:
|
||||
device = torch.device("cpu")
|
||||
|
||||
# Initialize distributed training
|
||||
if world_size > 1:
|
||||
dist.init_process_group(backend="nccl",
|
||||
init_method="env://",
|
||||
world_size=world_size,
|
||||
rank=rank)
|
||||
print(
|
||||
f"Initialized process: rank={rank}, local_rank={local_rank}, world_size={world_size}, device={device}"
|
||||
)
|
||||
|
||||
# Create dataset
|
||||
dataset = ParquetVideoTextDataset(
|
||||
args.path,
|
||||
batch_size=args.batch_size,
|
||||
rank=rank,
|
||||
world_size=world_size,
|
||||
)
|
||||
|
||||
# Create DataLoader with proper settings
|
||||
dataloader = StatefulDataLoader(
|
||||
dataset,
|
||||
batch_size=args.batch_size,
|
||||
num_workers=1, # Reduce number of workers to avoid memory issues
|
||||
prefetch_factor=2,
|
||||
shuffle=False,
|
||||
pin_memory=True,
|
||||
drop_last=True)
|
||||
|
||||
# Example of how to load dataloader state
|
||||
# if os.path.exists("/workspace/FastVideo/dataloader_state.pt"):
|
||||
# dataloader_state = torch.load("/workspace/FastVideo/dataloader_state.pt")
|
||||
# dataloader.load_state_dict(dataloader_state[rank])
|
||||
|
||||
# Warm-up with synchronization
|
||||
if rank == 0:
|
||||
print("Warming up...")
|
||||
for i, (latents, embeddings, masks, infos) in enumerate(dataloader):
|
||||
# Example of how to save dataloader state
|
||||
# if i == 30:
|
||||
# dist.barrier()
|
||||
# local_data = {rank: dataloader.state_dict()}
|
||||
# gathered_data = [None] * world_size
|
||||
# dist.all_gather_object(gathered_data, local_data)
|
||||
# if rank == 0:
|
||||
# global_state_dict = {}
|
||||
# for d in gathered_data:
|
||||
# global_state_dict.update(d)
|
||||
# torch.save(global_state_dict, "dataloader_state.pt")
|
||||
assert torch.sum(masks[0]).item() == torch.count_nonzero(
|
||||
embeddings[0]).item() // 4096
|
||||
if args.vae_debug:
|
||||
from diffusers.utils import export_to_video
|
||||
from diffusers.video_processor import VideoProcessor
|
||||
|
||||
from fastvideo.v1.configs.models.vaes import WanVAEConfig
|
||||
from fastvideo.v1.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.v1.models.loader.component_loader import VAELoader
|
||||
VAE_PATH = "/workspace/data/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/vae"
|
||||
fastvideo_args = FastVideoArgs(
|
||||
model_path=VAE_PATH,
|
||||
vae_config=WanVAEConfig(load_encoder=False),
|
||||
vae_precision="fp32")
|
||||
fastvideo_args.device = device
|
||||
vae_loader = VAELoader()
|
||||
vae = vae_loader.load(model_path=VAE_PATH,
|
||||
architecture="",
|
||||
fastvideo_args=fastvideo_args)
|
||||
|
||||
videoprocessor = VideoProcessor(vae_scale_factor=8)
|
||||
|
||||
with torch.inference_mode():
|
||||
video = vae.decode(latents[0].unsqueeze(0).to(device))
|
||||
video = videoprocessor.postprocess_video(video)
|
||||
video_path = os.path.join("/workspace/FastVideo/debug_videos",
|
||||
infos["caption"][0][:50] + ".mp4")
|
||||
export_to_video(video[0], video_path, fps=16)
|
||||
|
||||
# Move data to device
|
||||
# latents = latents.to(device)
|
||||
# embeddings = embeddings.to(device)
|
||||
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
# Benchmark
|
||||
if rank == 0:
|
||||
print(f"Benchmarking with batch_size={args.batch_size}")
|
||||
start_time = time.time()
|
||||
total_samples = 0
|
||||
for i, (latents, embeddings, masks,
|
||||
infos) in enumerate(tqdm.tqdm(dataloader, total=args.num_batches)):
|
||||
if i >= args.num_batches:
|
||||
break
|
||||
|
||||
# Move data to device
|
||||
latents = latents.to(device)
|
||||
embeddings = embeddings.to(device)
|
||||
|
||||
# Calculate actual batch size
|
||||
batch_size = latents.size(0)
|
||||
total_samples += batch_size
|
||||
|
||||
# Print progress only from rank 0
|
||||
if rank == 0 and (i + 1) % 10 == 0:
|
||||
elapsed = time.time() - start_time
|
||||
samples_per_sec = total_samples / elapsed
|
||||
print(
|
||||
f"Batch {i+1}/{args.num_batches}, Speed: {samples_per_sec:.2f} samples/sec"
|
||||
)
|
||||
|
||||
# Final statistics
|
||||
if world_size > 1:
|
||||
dist.barrier()
|
||||
|
||||
if rank == 0:
|
||||
elapsed = time.time() - start_time
|
||||
samples_per_sec = total_samples / elapsed
|
||||
|
||||
print("\nBenchmark Results:")
|
||||
print(f"Total time: {elapsed:.2f} seconds")
|
||||
print(f"Total samples: {total_samples}")
|
||||
print(f"Average speed: {samples_per_sec:.2f} samples/sec")
|
||||
print(f"Time per batch: {elapsed/args.num_batches*1000:.2f} ms")
|
||||
|
||||
if world_size > 1:
|
||||
dist.destroy_process_group()
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
@@ -138,6 +139,7 @@ class T2V_dataset(Dataset):
|
||||
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]
|
||||
@@ -270,7 +272,8 @@ class T2V_dataset(Dataset):
|
||||
cnt_resolution_mismatch += 1
|
||||
continue
|
||||
|
||||
# import ipdb;ipdb.set_trace()
|
||||
# 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 * (
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import random
|
||||
|
||||
import torch
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from huggingface_hub import HfApi, upload_folder
|
||||
|
||||
api = HfApi()
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
from typing import Any, Dict, List
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
|
||||
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
|
||||
"""
|
||||
Pad or crop an embedding [L, D] to exactly padding_length tokens.
|
||||
Return:
|
||||
- [L, D] tensor in pinned CPU memory
|
||||
- [L] attention mask in pinned CPU memory
|
||||
"""
|
||||
L, D = t.shape
|
||||
if padding_length > L: # pad
|
||||
pad = torch.zeros(padding_length - L, D, dtype=t.dtype, device=t.device)
|
||||
return torch.cat([t, pad], 0), torch.cat(
|
||||
[torch.ones(L), torch.zeros(padding_length - L)], 0)
|
||||
else: # crop
|
||||
return t[:padding_length], torch.ones(padding_length)
|
||||
|
||||
|
||||
def get_torch_tensors_from_row_dict(row_dict, keys) -> Dict[str, Any]:
|
||||
"""
|
||||
Get the latents and prompts from a row dictionary.
|
||||
"""
|
||||
return_dict = {}
|
||||
for key in keys:
|
||||
shape, bytes = None, None
|
||||
if isinstance(key, tuple):
|
||||
for k in key:
|
||||
try:
|
||||
shape = row_dict[f"{k}_shape"]
|
||||
bytes = row_dict[f"{k}_bytes"]
|
||||
except KeyError:
|
||||
continue
|
||||
key = key[0]
|
||||
if shape is None or bytes is None:
|
||||
raise ValueError(f"Key {key} not found in row_dict")
|
||||
else:
|
||||
shape = row_dict[f"{key}_shape"]
|
||||
bytes = row_dict[f"{key}_bytes"]
|
||||
|
||||
# TODO (peiyuan): read precision
|
||||
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]]:
|
||||
# Initialize tensors to hold padded embeddings and masks
|
||||
all_latents = []
|
||||
all_embs = []
|
||||
all_masks = []
|
||||
caption_text = []
|
||||
# Process each row individually
|
||||
for i, row in enumerate(batch_to_process):
|
||||
# Get tensors from row
|
||||
data = get_torch_tensors_from_row_dict(row, keys)
|
||||
latents, emb = data["vae_latent"], data["text_embedding"]
|
||||
|
||||
padded_emb, mask = pad(emb, text_padding_length)
|
||||
# Store in batch tensors
|
||||
all_latents.append(latents)
|
||||
all_embs.append(padded_emb)
|
||||
all_masks.append(mask)
|
||||
# TODO(py): remove this once we fix preprocess
|
||||
try:
|
||||
caption_text.append(row["prompt"])
|
||||
except KeyError:
|
||||
caption_text.append(row["caption"])
|
||||
|
||||
# Pin memory for faster transfer to GPU
|
||||
all_latents = torch.stack(all_latents)
|
||||
all_embs = torch.stack(all_embs)
|
||||
all_masks = torch.stack(all_masks)
|
||||
|
||||
return all_latents, all_embs, all_masks, caption_text
|
||||
@@ -2,21 +2,43 @@
|
||||
|
||||
from fastvideo.v1.distributed.communication_op import *
|
||||
from fastvideo.v1.distributed.parallel_state import (
|
||||
cleanup_dist_env_and_memory, get_sequence_model_parallel_rank,
|
||||
get_sequence_model_parallel_world_size, get_tensor_model_parallel_rank,
|
||||
get_tensor_model_parallel_world_size, get_world_group,
|
||||
init_distributed_environment, initialize_model_parallel,
|
||||
cleanup_dist_env_and_memory, get_dp_group, get_dp_rank, get_dp_world_size,
|
||||
get_sp_group, get_sp_parallel_rank, get_sp_world_size, get_torch_device,
|
||||
get_tp_group, get_tp_rank, get_tp_world_size, get_world_group,
|
||||
get_world_rank, get_world_size, init_distributed_environment,
|
||||
initialize_model_parallel,
|
||||
maybe_init_distributed_environment_and_model_parallel,
|
||||
model_parallel_is_initialized)
|
||||
from fastvideo.v1.distributed.utils import *
|
||||
|
||||
__all__ = [
|
||||
# Initialization
|
||||
"init_distributed_environment",
|
||||
"initialize_model_parallel",
|
||||
"get_sequence_model_parallel_rank",
|
||||
"get_sequence_model_parallel_world_size",
|
||||
"get_tensor_model_parallel_rank",
|
||||
"get_tensor_model_parallel_world_size",
|
||||
"cleanup_dist_env_and_memory",
|
||||
"get_world_group",
|
||||
"model_parallel_is_initialized",
|
||||
"maybe_init_distributed_environment_and_model_parallel",
|
||||
|
||||
# World group
|
||||
"get_world_group",
|
||||
"get_world_rank",
|
||||
"get_world_size",
|
||||
|
||||
# Data parallel group
|
||||
"get_dp_group",
|
||||
"get_dp_rank",
|
||||
"get_dp_world_size",
|
||||
|
||||
# Sequence parallel group
|
||||
"get_sp_group",
|
||||
"get_sp_parallel_rank",
|
||||
"get_sp_world_size",
|
||||
|
||||
# Tensor parallel group
|
||||
"get_tp_group",
|
||||
"get_tp_rank",
|
||||
"get_tp_world_size",
|
||||
|
||||
# Get torch device
|
||||
"get_torch_device",
|
||||
]
|
||||
|
||||
@@ -24,6 +24,7 @@ If you only need to use the distributed environment without model parallelism,
|
||||
"""
|
||||
import contextlib
|
||||
import gc
|
||||
import os
|
||||
import pickle
|
||||
import weakref
|
||||
from collections import namedtuple
|
||||
@@ -704,7 +705,7 @@ def init_world_group(ranks: List[int], local_rank: int,
|
||||
group_ranks=[ranks],
|
||||
local_rank=local_rank,
|
||||
torch_distributed_backend=backend,
|
||||
use_device_communicator=False,
|
||||
use_device_communicator=True,
|
||||
group_name="world",
|
||||
)
|
||||
|
||||
@@ -735,9 +736,6 @@ def get_tp_group() -> GroupCoordinator:
|
||||
return _TP
|
||||
|
||||
|
||||
# kept for backward compatibility
|
||||
get_tensor_model_parallel_group = get_tp_group
|
||||
|
||||
_ENABLE_CUSTOM_ALL_REDUCE = True
|
||||
|
||||
|
||||
@@ -747,10 +745,10 @@ def set_custom_all_reduce(enable: bool):
|
||||
|
||||
|
||||
def init_distributed_environment(
|
||||
world_size: int = -1,
|
||||
rank: int = -1,
|
||||
world_size: int = 1,
|
||||
rank: int = 0,
|
||||
distributed_init_method: str = "env://",
|
||||
local_rank: int = -1,
|
||||
local_rank: int = 0,
|
||||
backend: str = "nccl",
|
||||
):
|
||||
logger.debug(
|
||||
@@ -794,6 +792,14 @@ def get_sp_group() -> GroupCoordinator:
|
||||
return _SP
|
||||
|
||||
|
||||
_DP: Optional[GroupCoordinator] = None
|
||||
|
||||
|
||||
def get_dp_group() -> GroupCoordinator:
|
||||
assert _DP is not None, ("data parallel group is not initialized")
|
||||
return _DP
|
||||
|
||||
|
||||
def initialize_model_parallel(
|
||||
tensor_model_parallel_size: int = 1,
|
||||
sequence_model_parallel_size: int = 1,
|
||||
@@ -804,13 +810,13 @@ def initialize_model_parallel(
|
||||
|
||||
Arguments:
|
||||
tensor_model_parallel_size: number of GPUs used for tensor model
|
||||
parallelism.
|
||||
parallelism (used for language encoder).
|
||||
sequence_model_parallel_size: number of GPUs used for sequence model
|
||||
parallelism.
|
||||
parallelism (used for DiT).
|
||||
"""
|
||||
# Get world size and rank. Ensure some consistencies.
|
||||
assert torch.distributed.is_initialized()
|
||||
world_size: int = torch.distributed.get_world_size()
|
||||
assert _WORLD is not None, "world group is not initialized, please call init_distributed_environment first"
|
||||
world_size: int = get_world_size()
|
||||
backend = backend or torch.distributed.get_backend(
|
||||
get_world_group().device_group)
|
||||
|
||||
@@ -852,50 +858,83 @@ def initialize_model_parallel(
|
||||
backend,
|
||||
group_name="sp")
|
||||
|
||||
# Build the data parallel groups.
|
||||
num_data_parallel_groups: int = sequence_model_parallel_size
|
||||
global _DP
|
||||
assert _DP is None, ("data parallel group is already initialized")
|
||||
group_ranks = []
|
||||
|
||||
def get_sequence_model_parallel_world_size() -> int:
|
||||
for i in range(num_data_parallel_groups):
|
||||
ranks = list(range(i, world_size, num_data_parallel_groups))
|
||||
group_ranks.append(ranks)
|
||||
|
||||
_DP = init_model_parallel_group(group_ranks,
|
||||
get_world_group().local_rank,
|
||||
backend,
|
||||
group_name="dp")
|
||||
|
||||
|
||||
def get_sp_world_size() -> int:
|
||||
"""Return world size for the sequence model parallel group."""
|
||||
return get_sp_group().world_size
|
||||
|
||||
|
||||
def get_sequence_model_parallel_rank() -> int:
|
||||
def get_sp_parallel_rank() -> int:
|
||||
"""Return my rank for the sequence model parallel group."""
|
||||
return get_sp_group().rank_in_group
|
||||
|
||||
|
||||
def ensure_model_parallel_initialized(
|
||||
tensor_model_parallel_size: int,
|
||||
sequence_model_parallel_size: int,
|
||||
backend: Optional[str] = None,
|
||||
) -> None:
|
||||
"""Helper to initialize model parallel groups if they are not initialized,
|
||||
or ensure tensor-parallel, sequence-parallel sizes
|
||||
are equal to expected values if the model parallel groups are initialized.
|
||||
"""
|
||||
backend = backend or torch.distributed.get_backend(
|
||||
get_world_group().device_group)
|
||||
if not model_parallel_is_initialized():
|
||||
initialize_model_parallel(tensor_model_parallel_size,
|
||||
sequence_model_parallel_size, backend)
|
||||
def get_world_size() -> int:
|
||||
"""Return world size for the world group."""
|
||||
return get_world_group().world_size
|
||||
|
||||
|
||||
def get_world_rank() -> int:
|
||||
"""Return my rank for the world group."""
|
||||
return get_world_group().rank
|
||||
|
||||
|
||||
def get_dp_world_size() -> int:
|
||||
"""Return world size for the data parallel group."""
|
||||
return get_dp_group().world_size
|
||||
|
||||
|
||||
def get_dp_rank() -> int:
|
||||
"""Return my rank for the data parallel group."""
|
||||
return get_dp_group().rank_in_group
|
||||
|
||||
|
||||
def get_torch_device() -> torch.device:
|
||||
"""Return the torch device for the current rank."""
|
||||
return torch.device(f"cuda:{envs.LOCAL_RANK}")
|
||||
|
||||
|
||||
def maybe_init_distributed_environment_and_model_parallel(
|
||||
tp_size: int, sp_size: int, distributed_init_method: str = "env://"):
|
||||
if _WORLD is not None and model_parallel_is_initialized():
|
||||
# make sure the tp and sp sizes are correct
|
||||
assert get_tp_world_size(
|
||||
) == tp_size, f"You are trying to initialize model parallel groups with size {tp_size}, but they are already initialized with size {get_tp_world_size()}"
|
||||
assert get_sp_world_size(
|
||||
) == sp_size, f"You are trying to initialize model parallel groups with size {sp_size}, but they are already initialized with size {get_sp_world_size()}"
|
||||
return
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", 0))
|
||||
world_size = int(os.environ.get("WORLD_SIZE", 1))
|
||||
rank = int(os.environ.get("RANK", 0))
|
||||
|
||||
assert (
|
||||
get_tensor_model_parallel_world_size() == tensor_model_parallel_size
|
||||
), ("tensor parallel group already initialized, but of unexpected size: "
|
||||
f"{get_tensor_model_parallel_world_size()=} vs. "
|
||||
f"{tensor_model_parallel_size=}")
|
||||
|
||||
if sequence_model_parallel_size > 1:
|
||||
sp_world_size = get_sp_group().world_size
|
||||
assert (sp_world_size == sequence_model_parallel_size), (
|
||||
"sequence parallel group already initialized, but of unexpected size: "
|
||||
f"{sp_world_size=} vs. "
|
||||
f"{sequence_model_parallel_size=}")
|
||||
torch.cuda.set_device(local_rank)
|
||||
init_distributed_environment(
|
||||
world_size=world_size,
|
||||
rank=rank,
|
||||
local_rank=local_rank,
|
||||
distributed_init_method=distributed_init_method)
|
||||
initialize_model_parallel(tensor_model_parallel_size=tp_size,
|
||||
sequence_model_parallel_size=sp_size)
|
||||
|
||||
|
||||
def model_parallel_is_initialized() -> bool:
|
||||
"""Check if tensor, sequence parallel groups are initialized."""
|
||||
return _TP is not None and _SP is not None
|
||||
return _TP is not None and _SP is not None and _DP is not None
|
||||
|
||||
|
||||
_TP_STATE_PATCHED = False
|
||||
@@ -926,12 +965,12 @@ def patch_tensor_parallel_group(tp_group: GroupCoordinator):
|
||||
_TP = old_tp_group
|
||||
|
||||
|
||||
def get_tensor_model_parallel_world_size() -> int:
|
||||
def get_tp_world_size() -> int:
|
||||
"""Return world size for the tensor model parallel group."""
|
||||
return get_tp_group().world_size
|
||||
|
||||
|
||||
def get_tensor_model_parallel_rank() -> int:
|
||||
def get_tp_rank() -> int:
|
||||
"""Return my rank for the tensor model parallel group."""
|
||||
return get_tp_group().rank_in_group
|
||||
|
||||
@@ -948,6 +987,11 @@ def destroy_model_parallel() -> None:
|
||||
_SP.destroy()
|
||||
_SP = None
|
||||
|
||||
global _DP
|
||||
if _DP:
|
||||
_DP.destroy()
|
||||
_DP = None
|
||||
|
||||
|
||||
def destroy_distributed_environment() -> None:
|
||||
global _WORLD
|
||||
|
||||
@@ -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
|
||||
@@ -44,8 +42,8 @@ class VideoGenerator:
|
||||
Initialize the video generator.
|
||||
|
||||
Args:
|
||||
pipeline: The pipeline to use for inference
|
||||
fastvideo_args: The inference arguments
|
||||
executor_class: The executor class to use for inference
|
||||
"""
|
||||
self.fastvideo_args = fastvideo_args
|
||||
self.executor = executor_class(fastvideo_args)
|
||||
@@ -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,39 +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,
|
||||
device_str=device or "cuda" if torch.cuda.is_available() else "cpu",
|
||||
**config_args)
|
||||
fastvideo_args.check_fastvideo_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)
|
||||
|
||||
@@ -118,7 +91,6 @@ class VideoGenerator:
|
||||
# initialize_distributed_and_parallelism(fastvideo_args)
|
||||
|
||||
executor_class = Executor.get_class(fastvideo_args)
|
||||
|
||||
return cls(
|
||||
fastvideo_args=fastvideo_args,
|
||||
executor_class=executor_class,
|
||||
@@ -155,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)
|
||||
|
||||
@@ -181,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:
|
||||
@@ -243,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={},
|
||||
)
|
||||
|
||||
@@ -276,10 +249,10 @@ class VideoGenerator:
|
||||
|
||||
# Save video if requested
|
||||
if batch.save_video:
|
||||
save_path = batch.output_path
|
||||
if save_path:
|
||||
os.makedirs(os.path.dirname(save_path), exist_ok=True)
|
||||
video_path = os.path.join(save_path, f"{prompt[:100]}.mp4")
|
||||
output_path = batch.output_path
|
||||
if output_path:
|
||||
os.makedirs(output_path, exist_ok=True)
|
||||
video_path = os.path.join(output_path, f"{prompt[:100]}.mp4")
|
||||
imageio.mimsave(video_path, frames, fps=batch.fps, format="mp4")
|
||||
logger.info("Saved video to %s", video_path)
|
||||
else:
|
||||
@@ -295,6 +268,9 @@ class VideoGenerator:
|
||||
"generation_time": gen_time
|
||||
}
|
||||
|
||||
def set_lora_adapter(self, lora_nickname: str, lora_path: str) -> None:
|
||||
self.executor.set_lora_adapter(lora_nickname, lora_path)
|
||||
|
||||
def shutdown(self):
|
||||
"""
|
||||
Shutdown the video generator.
|
||||
|
||||
+133
-199
@@ -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
|
||||
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
|
||||
@@ -42,70 +48,38 @@ class FastVideoArgs:
|
||||
|
||||
# Parallelism
|
||||
num_gpus: int = 1
|
||||
tp_size: Optional[int] = None
|
||||
sp_size: Optional[int] = None
|
||||
tp_size: int = -1
|
||||
sp_size: int = -1
|
||||
hsdp_replicate_dim: int = 1
|
||||
hsdp_shard_dim: int = -1
|
||||
dist_timeout: Optional[int] = None # timeout for torch.distributed
|
||||
|
||||
# Video generation parameters
|
||||
embedded_cfg_scale: float = 6.0
|
||||
flow_shift: Optional[float] = None
|
||||
pipeline_config: PipelineConfig = field(default_factory=PipelineConfig)
|
||||
|
||||
output_type: str = "pil"
|
||||
|
||||
# DiT configuration
|
||||
dit_config: DiTConfig = field(default_factory=DiTConfig)
|
||||
precision: str = "bf16"
|
||||
use_cpu_offload: bool = True
|
||||
use_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 (Spatial-Temporal Attention) parameters
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
mask_strategy_file_path: Optional[str] = None
|
||||
STA_mode: Optional[str] = None
|
||||
skip_time_steps: int = 15
|
||||
|
||||
# Compilation
|
||||
enable_torch_compile: bool = False
|
||||
|
||||
use_cpu_offload: bool = False
|
||||
disable_autocast: bool = False
|
||||
|
||||
# StepVideo specific parameters
|
||||
pos_magic: Optional[str] = None
|
||||
neg_magic: Optional[str] = None
|
||||
timesteps_scale: Optional[bool] = None
|
||||
|
||||
# Logging
|
||||
log_level: str = "info"
|
||||
|
||||
# Inference parameters
|
||||
device_str: Optional[str] = None
|
||||
device = None
|
||||
# VSA parameters
|
||||
VSA_sparsity: float = 0.0 # inference/validation sparsity
|
||||
|
||||
@property
|
||||
def training_mode(self) -> bool:
|
||||
return not self.inference_mode
|
||||
|
||||
def __post_init__(self):
|
||||
pass
|
||||
self.check_fastvideo_args()
|
||||
|
||||
@staticmethod
|
||||
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
|
||||
@@ -116,11 +90,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,
|
||||
@@ -166,19 +135,29 @@ class FastVideoArgs:
|
||||
help="The number of GPUs to use.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--tensor-parallel-size",
|
||||
"--tp-size",
|
||||
type=int,
|
||||
default=FastVideoArgs.tp_size,
|
||||
help="The tensor parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--sequence-parallel-size",
|
||||
"--sp-size",
|
||||
type=int,
|
||||
default=FastVideoArgs.sp_size,
|
||||
help="The sequence parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--hsdp-replicate-dim",
|
||||
type=int,
|
||||
default=FastVideoArgs.hsdp_replicate_dim,
|
||||
help="The data parallelism size.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--hsdp-shard-dim",
|
||||
type=int,
|
||||
default=FastVideoArgs.hsdp_shard_dim,
|
||||
help="The data parallelism shards.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--dist-timeout",
|
||||
type=int,
|
||||
@@ -186,19 +165,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,
|
||||
@@ -207,53 +174,23 @@ class FastVideoArgs:
|
||||
help="Output type for the generated video",
|
||||
)
|
||||
|
||||
# STA (Sliding Tile Attention) parameters
|
||||
parser.add_argument(
|
||||
"--precision",
|
||||
"--STA-mode",
|
||||
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",
|
||||
default=FastVideoArgs.STA_mode,
|
||||
choices=[
|
||||
"STA_inference", "STA_searching", "STA_tuning",
|
||||
"STA_tuning_cfg", None
|
||||
],
|
||||
help="STA mode",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vae-tiling",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.vae_tiling,
|
||||
help="Enable VAE tiling",
|
||||
"--skip-time-steps",
|
||||
type=int,
|
||||
default=FastVideoArgs.skip_time_steps,
|
||||
help="Number of time steps to warmup (full attention) for STA",
|
||||
)
|
||||
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 (Spatial-Temporal Attention) parameters
|
||||
parser.add_argument(
|
||||
"--mask-strategy-file-path",
|
||||
type=str,
|
||||
@@ -269,8 +206,16 @@ class FastVideoArgs:
|
||||
parser.add_argument(
|
||||
"--use-cpu-offload",
|
||||
action=StoreBoolean,
|
||||
help="Use CPU offload for the model load",
|
||||
help=
|
||||
"Use CPU offload for model inference. Enable if run out of memory with FSDP.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--use-fsdp-inference",
|
||||
action=StoreBoolean,
|
||||
help=
|
||||
"Use FSDP for inference by sharding the model weights. Latency is very low due to prefetch--enable if run out of memory.",
|
||||
)
|
||||
|
||||
parser.add_argument(
|
||||
"--disable-autocast",
|
||||
action=StoreBoolean,
|
||||
@@ -278,75 +223,61 @@ 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
|
||||
parser.add_argument(
|
||||
"--log-level",
|
||||
type=str,
|
||||
default=FastVideoArgs.log_level,
|
||||
help="The logging level of all loggers.",
|
||||
)
|
||||
|
||||
# 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)
|
||||
kwargs[attr] = getattr(args, attr, default_value)
|
||||
value = getattr(args, attr, default_value)
|
||||
kwargs[attr] = value # type: ignore
|
||||
|
||||
return cls(**kwargs) # type: ignore
|
||||
|
||||
@classmethod
|
||||
def from_kwargs(cls, kwargs: Dict[str, Any]) -> "FastVideoArgs":
|
||||
kwargs['pipeline_config'] = PipelineConfig.from_kwargs(kwargs)
|
||||
return cls(**kwargs)
|
||||
|
||||
def check_fastvideo_args(self) -> None:
|
||||
"""Validate inference arguments for consistency"""
|
||||
if self.tp_size is None:
|
||||
if not self.inference_mode:
|
||||
assert self.hsdp_replicate_dim != -1, "hsdp_replicate_dim must be set for training"
|
||||
assert self.hsdp_shard_dim != -1, "hsdp_shard_dim must be set for training"
|
||||
assert self.sp_size != -1, "sp_size must be set for training"
|
||||
|
||||
if self.tp_size == -1:
|
||||
self.tp_size = self.num_gpus
|
||||
if self.sp_size is None:
|
||||
if self.sp_size == -1:
|
||||
self.sp_size = self.num_gpus
|
||||
if self.hsdp_shard_dim == -1:
|
||||
self.hsdp_shard_dim = self.num_gpus
|
||||
|
||||
assert self.sp_size <= self.num_gpus and self.num_gpus % self.sp_size == 0, "num_gpus must >= and be divisible by sp_size"
|
||||
assert self.hsdp_replicate_dim <= self.num_gpus and self.num_gpus % self.hsdp_replicate_dim == 0, "num_gpus must >= and be divisible by hsdp_replicate_dim"
|
||||
assert self.hsdp_shard_dim <= self.num_gpus and self.num_gpus % self.hsdp_shard_dim == 0, "num_gpus must >= and be divisible by hsdp_shard_dim"
|
||||
|
||||
if self.num_gpus < max(self.tp_size, self.sp_size):
|
||||
self.num_gpus = max(self.tp_size, self.sp_size)
|
||||
@@ -356,33 +287,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
|
||||
|
||||
@@ -402,7 +317,6 @@ def prepare_fastvideo_args(argv: List[str]) -> FastVideoArgs:
|
||||
FastVideoArgs.add_cli_args(parser)
|
||||
raw_args = parser.parse_args(argv)
|
||||
fastvideo_args = FastVideoArgs.from_cli_args(raw_args)
|
||||
fastvideo_args.check_fastvideo_args()
|
||||
global _current_fastvideo_args
|
||||
_current_fastvideo_args = fastvideo_args
|
||||
return fastvideo_args
|
||||
@@ -457,7 +371,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
|
||||
@@ -472,14 +385,14 @@ class TrainingArgs(FastVideoArgs):
|
||||
validation_steps: float = 0.0
|
||||
log_validation: bool = False
|
||||
tracker_project_name: str = ""
|
||||
# seed: int
|
||||
wandb_run_name: str = ""
|
||||
seed: Optional[int] = None
|
||||
|
||||
# output
|
||||
output_dir: str = ""
|
||||
checkpoints_total_limit: int = 0
|
||||
checkpointing_steps: int = 0
|
||||
resume_from_checkpoint: bool = False
|
||||
logging_dir: str = ""
|
||||
|
||||
# optimizer & scheduler
|
||||
num_train_epochs: int = 0
|
||||
@@ -487,7 +400,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
gradient_accumulation_steps: int = 0
|
||||
learning_rate: float = 0.0
|
||||
scale_lr: bool = False
|
||||
lr_scheduler: str = ""
|
||||
lr_scheduler: str = "constant"
|
||||
lr_warmup_steps: int = 0
|
||||
max_grad_norm: float = 0.0
|
||||
gradient_checkpointing: bool = False
|
||||
@@ -520,27 +433,29 @@ class TrainingArgs(FastVideoArgs):
|
||||
# master_weight_type
|
||||
master_weight_type: str = ""
|
||||
|
||||
# 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)
|
||||
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:
|
||||
@@ -630,6 +545,13 @@ 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,
|
||||
help="Seed for deterministic training")
|
||||
|
||||
# Output configuration
|
||||
parser.add_argument("--output-dir",
|
||||
@@ -764,4 +686,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
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user