Compare commits

...
Author SHA1 Message Date
SolitaryThinker e26b389f37 i2v validation 2025-06-11 21:35:10 -07:00
JerryZhou54 86dc4c4bb3 Small change 2025-06-11 13:29:11 -07:00
JerryZhou54 41d0400832 Fix preprocess 2025-06-11 13:29:11 -07:00
Wei Zhou e97aa17d33 Update preprocess_pipeline_i2v.py 2025-06-11 13:29:11 -07:00
JerryZhou54 0a40b79b36 I2V Runnable 2025-06-11 13:29:11 -07:00
“BrianChen1129” f8c69045d6 Add Encoded first frame to processed dataset 2025-06-11 13:29:10 -07:00
Zhang Peiyuan 0f2bbe71ac [misc] rename dp_size to hdsp_replicate_dim (#491) 2025-06-10 16:36:56 -07:00
Yongqi ChenandJerryZhou54 2a46902ecb [Feature][VSA]Update STA publish workflow (#498)
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
2025-06-10 19:33:34 -04:00
Zhang Peiyuan 66012d3a4c [Feat][Dataloader] 1/n Refactor parquet map-style dataloader (#492) 2025-06-10 16:00:13 -07:00
William Lin f666b9de41 [misc] Add missing license headers (#499) 2025-06-10 14:25:32 -07:00
Yongqi Chen 7e3c073b55 [Feature] Adding VSA inference (#478) 2025-06-10 16:03:53 -04:00
Wei Zhou a6aa21bd07 [bugfix][Cli Inference] Resolve runtime errors when running fastvideo generate (#495) 2025-06-10 02:50:21 -04:00
Wenxuan Tan 6519b57aab [chore] Fix main pre-commit CI failure (#494) 2025-06-10 00:54:37 -05:00
Wei Zhou 675aea6ece [bugfix][Cli Inference] Resolve runtime errors when running fastvideo generate (#493) 2025-06-09 19:49:34 -07:00
Zhang Peiyuan 46e7a15e0d [misc] Improve distributed related env variables and setup (#487) 2025-06-08 09:14:48 -07:00
Yongqi ChenandJerryZhou54 e4f702d7ec [Bug] Fix multi gpus issues in v1 scripts (#489)
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
2025-06-07 21:55:49 -07:00
Wenxuan Tan bb68fcc809 Revert "Add torch.compile for all small ops" (#484) 2025-06-07 07:21:17 -05:00
Wenxuan Tan b392e6a874 Add torch.compile for all small ops (#432) 2025-06-06 21:10:42 -07:00
Zhang PeiyuanandWill Lin 0991003905 [bugfix] [misc] fix denoising stage init; rename distributed env function; fix logging. (#481)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-06-06 20:01:10 -07:00
147 changed files with 5394 additions and 1310 deletions
+17 -13
View File
@@ -5,7 +5,7 @@ on:
branches:
- main
paths:
- "csrc/sliding_tile_attention/setup.py"
- "csrc/attn/setup_sta.py"
workflow_dispatch:
jobs:
@@ -23,13 +23,13 @@ jobs:
- name: Check if version changed
id: check-version
run: |
cd csrc/sliding_tile_attention
cd csrc/attn
# Get current commit's version
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_sta.py)
echo "New version: $NEW_VERSION"
# Get previous version from git history
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
OLD_VERSION=$(git show HEAD~1:./setup_sta.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
@@ -136,19 +136,21 @@ jobs:
- name: Build wheel
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
# However this still fails so I'm using a newer version of setuptools
pip install setuptools
pip install ninja packaging wheel
cd csrc/sliding_tile_attention # Move into the correct folder
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
python setup.py bdist_wheel --dist-dir=dist
cd csrc/attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup_sta.py bdist_wheel --dist-dir=dist
- name: Rename wheel file
run: |
cd csrc/sliding_tile_attention
cd csrc/attn
CUDA_SHORT_VERSION=$(echo ${{ matrix.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-version }} | cut -d. -f1,2)
@@ -163,7 +165,7 @@ jobs:
uses: actions/upload-artifact@v4
with:
name: ${{ env.wheel_name }}
path: csrc/sliding_tile_attention/dist/*.whl
path: csrc/attn/dist/*.whl
retention-days: 90
publish_package:
@@ -229,17 +231,19 @@ jobs:
- name: Build source distribution
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
# However this still fails so I'm using a newer version of setuptools
pip install setuptools
pip install ninja packaging wheel
cd csrc/sliding_tile_attention # Move into the correct folder
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
python setup.py sdist --dist-dir=dist
cd csrc/attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup_sta.py sdist --dist-dir=dist
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: csrc/sliding_tile_attention/dist/
packages-dir: csrc/attn/dist/
+1 -1
View File
@@ -28,4 +28,4 @@ jobs:
- name: Run Pytest
run: |
pytest --ignore csrc/sliding_tile_attention/test
pytest --ignore csrc/attn/test
+2 -2
View File
@@ -1,3 +1,3 @@
[submodule "csrc/sliding_tile_attention/tk"]
path = csrc/sliding_tile_attention/tk
[submodule "csrc/attn/tk"]
path = csrc/attn/tk
url = https://github.com/HazyResearch/ThunderKittens.git
@@ -7,26 +7,26 @@
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
First, install C++20 for ThunderKittens:
```bash
sudo apt update
sudo apt install gcc-11 g++-11
sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave /usr/bin/g++ g++ /usr/bin/g++-11
sudo apt update
sudo apt install clang-11
```
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),
+225
View File
@@ -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'
+15
View File
@@ -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
+76
View File
@@ -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]
@@ -829,3 +829,4 @@ sta_forward(torch::Tensor q, torch::Tensor k, torch::Tensor v, torch::Tensor o,
return o;
cudaDeviceSynchronize();
}
+266
View File
@@ -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()
+136
View File
@@ -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.")
+175
View File
@@ -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']}")
+27
View File
@@ -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
}
+469
View File
@@ -0,0 +1,469 @@
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_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
@@ -0,0 +1,31 @@
{
"data": [
{
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
"image_path": null,
"video_path": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-034.mp4",
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-027.mp4",
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": "examples/training/finetune/wan_i2v_14b_480p/crush_smol/validation_dataset/yYcK4nANZz4-Scene-030.mp4",
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -104,13 +104,7 @@ if __name__ == "__main__":
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
args = parser.parse_args()
main(args)
@@ -7,7 +7,7 @@ import torch.distributed as dist
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import maybe_download_model, shallow_asdict
from fastvideo.v1.distributed import init_distributed_environment, initialize_model_parallel
from fastvideo.v1.distributed import maybe_init_distributed_environment_and_model_parallel, get_world_size
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo import PipelineConfig
@@ -18,15 +18,7 @@ logger = init_logger(__name__)
def main(args):
args.model_path = maybe_download_model(args.model_path)
# Assume using torchrun
local_rank = int(os.getenv("RANK", 0))
rank = int(os.environ.get("RANK", 0))
world_size = int(os.getenv("WORLD_SIZE", 1))
init_distributed_environment(world_size=world_size, rank=rank, local_rank=local_rank)
initialize_model_parallel(tensor_model_parallel_size=world_size, sequence_model_parallel_size=world_size)
torch.cuda.set_device(local_rank)
if not dist.is_initialized():
dist.init_process_group(backend="nccl", init_method="env://", world_size=world_size, rank=local_rank)
maybe_init_distributed_environment_and_model_parallel(1, 1)
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
@@ -37,12 +29,9 @@ def main(args):
pipeline_config_args = shallow_asdict(pipeline_config)
pipeline_config_args.update(kwargs)
fastvideo_args = FastVideoArgs(model_path=args.model_path,
num_gpus=world_size,
device_str="cuda",
num_gpus=get_world_size(),
**pipeline_config_args,
)
fastvideo_args.check_fastvideo_args()
fastvideo_args.device = torch.device(f"cuda:{local_rank}")
PreprocessPipeline = PreprocessPipeline_I2V if args.preprocess_task == "i2v" else PreprocessPipeline_T2V
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
@@ -54,7 +43,7 @@ if __name__ == "__main__":
parser.add_argument("--model_path", type=str, default="data/mochi")
parser.add_argument("--model_type", type=str, default="mochi")
parser.add_argument("--data_merge_path", type=str, required=True)
parser.add_argument("--validation_prompt_txt", type=str)
parser.add_argument("--validation_dataset_file", type=str)
parser.add_argument("--num_frames", type=int, default=163)
parser.add_argument(
"--dataloader_num_workers",
@@ -108,13 +97,6 @@ if __name__ == "__main__":
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
args = parser.parse_args()
main(args)
-7
View File
@@ -671,13 +671,6 @@ if __name__ == "__main__":
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
-7
View File
@@ -693,13 +693,6 @@ if __name__ == "__main__":
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
+1 -7
View File
@@ -520,13 +520,7 @@ if __name__ == "__main__":
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
import json
import os
from collections import defaultdict
+4 -1
View File
@@ -3,11 +3,14 @@
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.v1.attention.layer import DistributedAttention, LocalAttention
from fastvideo.v1.attention.layer import (DistributedAttention,
DistributedAttention_VSA,
LocalAttention)
from fastvideo.v1.attention.selector import get_attn_backend
__all__ = [
"DistributedAttention",
"DistributedAttention_VSA",
"LocalAttention",
"AttentionBackend",
"AttentionMetadata",
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from typing import List, Optional, Type
import torch
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from typing import List, Optional, Type
import torch
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
import json
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Type
@@ -139,7 +140,7 @@ class SlidingTileAttentionImpl(AttentionImpl):
self.sp_size = sp_group.world_size
# STA config
self.STA_base_tile_size = [6, 8, 8]
self.img_latent_shape_mapping = RangeDict({
self.dit_seq_shape_mapping = RangeDict({
(115200, 115456): '30x48x80',
82944: '36x48x48',
69120: '18x48x80',
@@ -154,9 +155,9 @@ class SlidingTileAttentionImpl(AttentionImpl):
x = rearrange(x,
"b (sp t h w) head d -> b (t sp h w) head d",
sp=self.sp_size,
t=self.img_latent_shape_int[0] // self.sp_size,
h=self.img_latent_shape_int[1],
w=self.img_latent_shape_int[2])
t=self.dit_seq_shape_int[0] // self.sp_size,
h=self.dit_seq_shape_int[1],
w=self.dit_seq_shape_int[2])
return rearrange(
x,
"b (n_t ts_t n_h ts_h n_w ts_w) h d -> b (n_t n_h n_w ts_t ts_h ts_w) h d",
@@ -180,9 +181,9 @@ class SlidingTileAttentionImpl(AttentionImpl):
return rearrange(x,
"b (t sp h w) head d -> b (sp t h w) head d",
sp=self.sp_size,
t=self.img_latent_shape_int[0] // self.sp_size,
h=self.img_latent_shape_int[1],
w=self.img_latent_shape_int[2])
t=self.dit_seq_shape_int[0] // self.sp_size,
h=self.dit_seq_shape_int[1],
w=self.dit_seq_shape_int[2])
def preprocess_qkv(
self,
@@ -190,14 +191,12 @@ class SlidingTileAttentionImpl(AttentionImpl):
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
img_sequence_length = qkv.shape[1]
self.img_latent_shape_str = self.img_latent_shape_mapping[
img_sequence_length]
self.full_window_size = self.full_window_mapping[
self.img_latent_shape_str]
self.img_latent_shape_int = list(
map(int, self.img_latent_shape_str.split('x')))
self.img_seq_length = self.img_latent_shape_int[
0] * self.img_latent_shape_int[1] * self.img_latent_shape_int[2]
self.dit_seq_shape_str = self.dit_seq_shape_mapping[img_sequence_length]
self.full_window_size = self.full_window_mapping[self.dit_seq_shape_str]
self.dit_seq_shape_int = list(
map(int, self.dit_seq_shape_str.split('x')))
self.img_seq_length = self.dit_seq_shape_int[
0] * self.dit_seq_shape_int[1] * self.dit_seq_shape_int[2]
return self.tile(qkv)
def postprocess_output(
@@ -252,12 +251,12 @@ class SlidingTileAttentionImpl(AttentionImpl):
for window_size in STA_param[:-1]:
sparse_hidden_states = sliding_tile_attention(
query, key, value, [window_size] * head_num, text_length,
has_text, self.img_latent_shape_str).transpose(1, 2)
has_text, self.dit_seq_shape_str).transpose(1, 2)
sparse_attn_hidden_states_all.append(sparse_hidden_states)
hidden_states = sliding_tile_attention(
query, key, value, [full_mask_window] * head_num, text_length,
has_text, self.img_latent_shape_str).transpose(1, 2)
has_text, self.dit_seq_shape_str).transpose(1, 2)
attn_L2_loss = []
attn_L1_loss = []
@@ -288,18 +287,12 @@ class SlidingTileAttentionImpl(AttentionImpl):
forward_batch.mask_search_final_result_pos[timestep].append(
layer_loss_save)
else:
# windows = [
# self.mask_strategy[timestep][layer_idx][head_idx + start_head]
# for head_idx in range(head_num)
# ]
windows = [
STA_param[head_idx + start_head] for head_idx in range(head_num)
]
# if has_text is False:
# from IPython import embed
# embed()
hidden_states = sliding_tile_attention(
query, key, value, windows, text_length, has_text,
self.img_latent_shape_str).transpose(1, 2)
self.dit_seq_shape_str).transpose(1, 2)
return hidden_states
@@ -0,0 +1,185 @@
# SPDX-License-Identifier: Apache-2.0
import math
from dataclasses import dataclass
from typing import List, Optional, Type
import torch
from einops import rearrange
from vsa import video_sparse_attn
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.v1.distributed import get_sp_group
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
class VideoSparseAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_supported_head_sizes() -> List[int]:
return [64, 128]
@staticmethod
def get_name() -> str:
return "VIDEO_SPARSE_ATTN"
@staticmethod
def get_impl_cls() -> Type["VideoSparseAttentionImpl"]:
return VideoSparseAttentionImpl
@staticmethod
def get_metadata_cls() -> Type["VideoSparseAttentionMetadata"]:
return VideoSparseAttentionMetadata
@staticmethod
def get_builder_cls() -> Type["VideoSparseAttentionMetadataBuilder"]:
return VideoSparseAttentionMetadataBuilder
@dataclass
class VideoSparseAttentionMetadata(AttentionMetadata):
current_timestep: int
dit_seq_shape: List[int]
VSA_sparsity: float
class VideoSparseAttentionMetadataBuilder(AttentionMetadataBuilder):
def __init__(self):
pass
def prepare(self):
pass
def build(
self,
current_timestep: int,
forward_batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> VideoSparseAttentionMetadata:
if forward_batch.latents is None:
raise ValueError("latents cannot be None")
raw_latent_shape = forward_batch.latents.shape
patch_size = fastvideo_args.dit_config.patch_size
dit_seq_shape = [
raw_latent_shape[2] // patch_size[0],
raw_latent_shape[3] // patch_size[1],
raw_latent_shape[4] // patch_size[2]
]
VSA_sparsity = forward_batch.VSA_sparsity
return VideoSparseAttentionMetadata(current_timestep=current_timestep,
dit_seq_shape=dit_seq_shape,
VSA_sparsity=VSA_sparsity)
class VideoSparseAttentionImpl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
causal: bool,
softmax_scale: float,
num_kv_heads: Optional[int] = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.prefix = prefix
sp_group = get_sp_group()
self.sp_size = sp_group.world_size
self.VSA_base_tile_size = [4, 4, 4]
self.dit_seq_shape: List[int]
self.full_window_size: List[int]
self.img_seq_length: int
def tile(self, x: torch.Tensor) -> torch.Tensor:
x = rearrange(x,
"b (sp t h w) head d -> b (t sp h w) head d",
sp=self.sp_size,
t=self.dit_seq_shape[0] // self.sp_size,
h=self.dit_seq_shape[1],
w=self.dit_seq_shape[2])
return rearrange(
x,
"b (n_t ts_t n_h ts_h n_w ts_w) h d -> b (n_t n_h n_w ts_t ts_h ts_w) h d",
n_t=self.full_window_size[0],
n_h=self.full_window_size[1],
n_w=self.full_window_size[2],
ts_t=self.VSA_base_tile_size[0],
ts_h=self.VSA_base_tile_size[1],
ts_w=self.VSA_base_tile_size[2])
def untile(self, x: torch.Tensor) -> torch.Tensor:
x = rearrange(
x,
"b (n_t n_h n_w ts_t ts_h ts_w) h d -> b (n_t ts_t n_h ts_h n_w ts_w) h d",
n_t=self.full_window_size[0],
n_h=self.full_window_size[1],
n_w=self.full_window_size[2],
ts_t=self.VSA_base_tile_size[0],
ts_h=self.VSA_base_tile_size[1],
ts_w=self.VSA_base_tile_size[2])
return rearrange(x,
"b (t sp h w) head d -> b (sp t h w) head d",
sp=self.sp_size,
t=self.dit_seq_shape[0] // self.sp_size,
h=self.dit_seq_shape[1],
w=self.dit_seq_shape[2])
def preprocess_qkv(
self,
qkv: torch.Tensor,
attn_metadata: VideoSparseAttentionMetadata,
) -> torch.Tensor:
self.dit_seq_shape = attn_metadata.dit_seq_shape
self.full_window_size = [
self.dit_seq_shape[0] // self.VSA_base_tile_size[0],
self.dit_seq_shape[1] // self.VSA_base_tile_size[1],
self.dit_seq_shape[2] // self.VSA_base_tile_size[2]
]
self.img_seq_length = math.prod(self.dit_seq_shape)
return self.tile(qkv)
def postprocess_output(
self,
output: torch.Tensor,
attn_metadata: VideoSparseAttentionMetadata,
) -> torch.Tensor:
return self.untile(output)
def forward( # type: ignore[override]
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
gate_compress: torch.Tensor,
attn_metadata: VideoSparseAttentionMetadata,
) -> torch.Tensor:
query = query.transpose(1, 2).contiguous()
key = key.transpose(1, 2).contiguous()
value = value.transpose(1, 2).contiguous()
gate_compress = gate_compress.transpose(1, 2).contiguous()
cur_topk = math.ceil(
(1 - attn_metadata.VSA_sparsity) *
(self.img_seq_length / math.prod(self.VSA_base_tile_size)))
hidden_states = video_sparse_attn(
query,
key,
value,
topk=cur_topk,
block_size=(4, 4, 4),
compress_attn_weight=gate_compress).transpose(1, 2)
return hidden_states
+71 -4
View File
@@ -9,8 +9,8 @@ from fastvideo.v1.attention.selector import (backend_name_to_enum,
get_attn_backend)
from fastvideo.v1.distributed.communication_op import (
sequence_model_parallel_all_gather, sequence_model_parallel_all_to_all_4D)
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size)
from fastvideo.v1.distributed.parallel_state import (get_sp_parallel_rank,
get_sp_world_size)
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.utils import get_compute_dtype
@@ -86,8 +86,8 @@ class DistributedAttention(nn.Module):
assert q.dim() == 4 and k.dim() == 4 and v.dim(
) == 4, "Expected 4D tensors"
batch_size, seq_len, num_heads, head_dim = q.shape
local_rank = get_sequence_model_parallel_rank()
world_size = get_sequence_model_parallel_world_size()
local_rank = get_sp_parallel_rank()
world_size = get_sp_world_size()
forward_context: ForwardContext = get_forward_context()
ctx_attn_metadata = forward_context.attn_metadata
@@ -135,6 +135,73 @@ class DistributedAttention(nn.Module):
return output, replicated_output
class DistributedAttention_VSA(DistributedAttention):
"""Distributed attention layer with VSA support.
"""
def forward(
self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
replicated_q: Optional[torch.Tensor] = None,
replicated_k: Optional[torch.Tensor] = None,
replicated_v: Optional[torch.Tensor] = None,
gate_compress: Optional[torch.Tensor] = None,
) -> tuple[torch.Tensor, Optional[torch.Tensor]]:
"""Forward pass for distributed attention.
Args:
q (torch.Tensor): Query tensor [batch_size, seq_len, num_heads, head_dim]
k (torch.Tensor): Key tensor [batch_size, seq_len, num_heads, head_dim]
v (torch.Tensor): Value tensor [batch_size, seq_len, num_heads, head_dim]
gate_compress (torch.Tensor): Gate compress tensor [batch_size, seq_len, num_heads, head_dim]
replicated_q (Optional[torch.Tensor]): Replicated query tensor, typically for text tokens
replicated_k (Optional[torch.Tensor]): Replicated key tensor
replicated_v (Optional[torch.Tensor]): Replicated value tensor
Returns:
Tuple[torch.Tensor, Optional[torch.Tensor]]: A tuple containing:
- o (torch.Tensor): Output tensor after attention for the main sequence
- replicated_o (Optional[torch.Tensor]): Output tensor for replicated tokens, if provided
"""
# Check text tokens are not supported for VSA now
assert replicated_q is None and replicated_k is None and replicated_v is None, "Replicated QKV is not supported for VSA now"
# Check input shapes
assert q.dim() == 4 and k.dim() == 4 and v.dim(
) == 4, "Expected 4D tensors"
forward_context: ForwardContext = get_forward_context()
ctx_attn_metadata = forward_context.attn_metadata
# Stack QKV
qkvg = torch.cat([q, k, v, gate_compress],
dim=0) # [3, seq_len, num_heads, head_dim]
# Redistribute heads across sequence dimension
qkvg = sequence_model_parallel_all_to_all_4D(qkvg,
scatter_dim=2,
gather_dim=1)
qkvg = self.impl.preprocess_qkv(
qkvg, ctx_attn_metadata) # (yongqi) pass latent shape here?
q, k, v, gate_compress = qkvg.chunk(4, dim=0)
output = self.impl.forward(q, k, v, gate_compress,
ctx_attn_metadata) # type: ignore[call-arg]
# Redistribute back if using sequence parallelism
replicated_output = None
# Apply backend-specific postprocess_output
output = self.impl.postprocess_output(output, ctx_attn_metadata)
output = sequence_model_parallel_all_to_all_4D(output,
scatter_dim=1,
gather_dim=2)
return output, replicated_output
class LocalAttention(nn.Module):
"""Attention layer.
"""
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field, fields
from typing import Any, Dict
+3 -1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import Any, List, Optional, Tuple
@@ -16,7 +17,8 @@ class DiTArchConfig(ArchConfig):
...] = (_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
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import List, Optional, Tuple
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import List, Optional, Tuple, Union
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import List, Optional, Tuple
@@ -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
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import Any, Union
@@ -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
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
import json
from dataclasses import asdict, dataclass, field, fields
from typing import Any, Callable, Dict, Optional, Tuple, cast
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import Callable, Tuple, TypedDict
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
"""Registry for pipeline weight-specific configurations."""
import os
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.v1.configs.models import DiTConfig, VAEConfig
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from typing import Callable, Tuple
+8
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from typing import Any, Dict, List, Optional, Union
@@ -38,6 +39,7 @@ class SamplingParam:
num_inference_steps: int = 50
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
VSA_sparsity: float = 0.0
# TeaCache parameters
enable_teacache: bool = False
@@ -183,6 +185,12 @@ class SamplingParam:
default=SamplingParam.image_path,
help="Path to input image for image-to-video generation",
)
parser.add_argument(
"--VSA-sparsity",
type=float,
default=SamplingParam.VSA_sparsity,
help="VSA attention sparsity",
)
return parser
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.v1.configs.sample.base import SamplingParam
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
import os
from typing import Any, Callable, Dict, Optional
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.v1.configs.sample.base import SamplingParam
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.v1.configs.sample.base import CacheParams
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.v1.configs.sample.base import SamplingParam
+4
View File
@@ -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,170 @@
# 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/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 2 --verify_resume
'''
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
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:
for i, (latents, embeddings, masks,
data_indices) in enumerate(dataloader):
logger.info("Batch %d data_indices: %s", i, data_indices)
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()
dataloader = build_parquet_map_style_dataloader(args.path,
args.batch_size,
args.num_data_workers)
# Load dataloader state using distributed checkpoint
logger.info("Rank %d: Loading dataloader state from %s",
get_world_rank(), checkpoint_dir)
load_states = {"dataloader": dataloader}
dist_cp.load(load_states, checkpoint_id=checkpoint_dir.as_posix())
logger.info("Rank %d: Loaded dataloader state from %s",
get_world_rank(), checkpoint_dir)
for i, (latents, embeddings, masks,
data_indices) in enumerate(dataloader):
logger.info("Batch %d data_indices: %s", i, data_indices)
if i >= args.num_batches_per_epoch - 1:
break
logger.info("Restart from the beginning")
dataloader = build_parquet_map_style_dataloader(args.path,
args.batch_size,
args.num_data_workers)
for i, (latents, embeddings, masks,
data_indices) in enumerate(dataloader):
logger.info("Batch %d data_indices: %s", i, data_indices)
if i >= args.num_batches_per_epoch * 2 - 1:
break
start_time = time.time()
total_samples = 0
total_batches = 0
for _ in range(args.num_epoch):
for i, (latents, embeddings, masks,
data_indices) in enumerate(dataloader):
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
@@ -34,6 +35,9 @@ pyarrow_schema_i2v = pa.schema([
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_dtype", pa.string()),
pa.field("encoded_first_frame_bytes", pa.binary()),
pa.field("encoded_first_frame_shape", pa.list_(pa.int64())),
pa.field("encoded_first_frame_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import json
import os
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
import json
import os
import random
@@ -0,0 +1,323 @@
# SPDX-License-Identifier: Apache-2.0
import os
from typing import Any, Dict, List, Tuple
import numpy as np
import pyarrow.parquet as pq
# Torch in general
import torch
# Dataset
from torch.utils.data import Dataset, Sampler
from torchdata.stateful_dataloader import StatefulDataLoader
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:
# add more indices to make it divisible by (batch_size * num_sp_groups)
padding_size = self.num_sp_groups * self.batch_size - (
self.dataset_size % (self.num_sp_groups * self.batch_size))
global_indices = torch.cat(
[global_indices, global_indices[:padding_size]])
# shard the indices to each sp group
ith_sp_group = self.global_rank // self.sp_world_size
sp_group_local_indices = global_indices[ith_sp_group::self.
num_sp_groups]
self.sp_group_local_indices = sp_group_local_indices
logger.info("sp_group_local_indices: %d", len(sp_group_local_indices))
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):
lengths = []
file_names = []
for root, _, files in os.walk(path):
for file in sorted(files):
if file.endswith('.parquet'):
file_path = os.path.join(root, file)
num_rows = pq.ParquetFile(file_path).metadata.num_rows
lengths.append(num_rows)
file_names.append(file_path)
# sort according to file name to ensure all rank has the same order (in case os.walk is not sorted)
file_names_sorted, lengths_sorted = zip(
*sorted(zip(file_names, lengths), key=lambda x: x[0]))
assert len(file_names_sorted) != 0, "No parquet files found in the dataset"
return file_names_sorted, lengths_sorted
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", "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_torch_tensors_from_row_dict(
self, row_dict: Dict[str, Any]) -> Dict[str, torch.Tensor]:
"""
Get the latents and prompts from a row dictionary.
"""
return_dict = {}
for key in self.keys:
shape = row_dict[f"{key}_shape"]
bytes = row_dict[f"{key}_bytes"]
# TODO (peiyuan): read precision
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
data = torch.from_numpy(data)
return_dict[key] = data
return return_dict
def get_validation_negative_prompt(self) -> tuple[Any, Any, Any, Any]:
"""
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]])
# Get tensors using the existing helper method
data = self._get_torch_tensors_from_row_dict(row_dict)
emb = data["text_embedding"]
# Pad the embedding and get mask
padded_emb, mask = self._pad(emb, self.text_padding_length)
# Pin memory for faster transfer to GPU
padded_emb = padded_emb
mask = mask
return None, padded_emb, mask, None
def _pad(self, t: torch.Tensor, padding_length: int) -> torch.Tensor:
"""
Pad or crop an embedding [L, D] to exactly padding_length tokens.
Return:
- [L, D] tensor in pinned CPU memory
- [L] attention mask in pinned CPU memory
"""
L, D = t.shape
if padding_length > L: # pad
pad = torch.zeros(padding_length - L,
D,
dtype=t.dtype,
device=t.device)
return torch.cat([t, pad], 0), torch.cat(
[torch.ones(L), torch.zeros(padding_length - L)], 0)
else: # crop
return t[:padding_length], torch.ones(padding_length)
# 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
]
# Initialize tensors to hold padded embeddings and masks
all_latents = []
all_embs = []
all_masks = []
# Process each row individually
for i, row in enumerate(rows):
# Get tensors from row
data = self._get_torch_tensors_from_row_dict(row)
print(data)
import pdb; pdb.set_trace()
latents, emb = data["vae_latent"], data["text_embedding"]
padded_emb, mask = self._pad(emb, self.text_padding_length)
# Store in batch tensors
all_latents.append(latents)
all_embs.append(padded_emb)
all_masks.append(mask)
# Pin memory for faster transfer to GPU
all_latents = torch.stack(all_latents)
all_embs = torch.stack(all_embs)
all_masks = torch.stack(all_masks)
return all_latents, all_embs, all_masks, indices
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
-470
View File
@@ -1,470 +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_dp_group,
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,
validation: bool = False):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
self.rank = rank
self.local_rank = get_sequence_model_parallel_rank()
self.sp_group = get_sp_group()
self.dp_group = get_dp_group()
self.dp_world_size = self.dp_group.world_size
self.sp_world_size = self.sp_group.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.validation = validation
# Negative prompt caching
self.neg_metadata = None
self.cached_neg_prompt: Dict[str, Any] | None = None
self.plan_output_dir = os.path.join(
self.path,
f"data_plan_{self.world_size}_{self.sp_world_size}_{self.dp_world_size}.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}")
else:
print(f"Creating new plan for {self.plan_output_dir}")
# 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))
# the negative prompt is always the first row in the first
# parquet file
if validation:
self.neg_metadata = metadatas[0]
metadatas = metadatas[1:]
# 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)
if validation:
assert self.neg_metadata is not None
plan["negative_prompt"] = [self.neg_metadata]
with open(self.plan_output_dir, "w") as f:
json.dump(plan, f)
else:
pass
dist.barrier()
if validation:
with open(self.plan_output_dir) as f:
plan = json.load(f)
self.neg_metadata = plan["negative_prompt"][0]
def _load_and_cache_negative_prompt(self) -> None:
"""Load and cache the negative prompt. Only rank 0 in each SP group should call this."""
if not self.validation or self.neg_metadata is None:
return
if self.cached_neg_prompt is not None:
return
# Only rank 0 in each SP group should read the negative prompt
try:
file_path, row_idx = self.neg_metadata
parquet_file = pq.ParquetFile(file_path)
# Since negative prompt is always the first row (row_idx = 0),
# it's always in the first row group
row_group_index = 0
local_index = row_idx # This will be 0 for the negative prompt
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
# Process the negative prompt row
self.cached_neg_prompt = self._process_row(row_dict)
except Exception as e:
logger.error("Failed to load negative prompt: %s", e)
self.cached_neg_prompt = None
def get_validation_negative_prompt(
self
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, Dict[str, Any]]:
"""
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).
"""
if not self.validation:
raise ValueError(
"get_validation_negative_prompt() can only be called in validation mode"
)
# Load and cache if needed (only rank 0 in SP group will actually load)
if self.cached_neg_prompt is None:
self._load_and_cache_negative_prompt()
if self.cached_neg_prompt is None:
raise RuntimeError(
f"Rank {self.rank} (SP rank {self.local_rank}): Could not retrieve negative prompt data"
)
# Extract the components
lat, emb, mask, info = (self.cached_neg_prompt["latents"],
self.cached_neg_prompt["embeddings"],
self.cached_neg_prompt["masks"],
self.cached_neg_prompt["info"])
# Apply the same processing as in __getitem__
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 __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 > row_idx:
row_group_index = i
local_index = 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
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
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
import json
import math
import os
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
import random
import torch
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from huggingface_hub import HfApi, upload_folder
api = HfApi()
+31 -15
View File
@@ -2,27 +2,43 @@
from fastvideo.v1.distributed.communication_op import *
from fastvideo.v1.distributed.parallel_state import (
cleanup_dist_env_and_memory, get_data_parallel_rank,
get_data_parallel_world_size, get_dp_group,
get_sequence_model_parallel_rank, get_sequence_model_parallel_world_size,
get_sp_group, 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_data_parallel_world_size",
"get_data_parallel_rank",
"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",
"get_dp_group",
"get_sp_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",
]
+47 -45
View File
@@ -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
@@ -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
@@ -805,7 +803,6 @@ def get_dp_group() -> GroupCoordinator:
def initialize_model_parallel(
tensor_model_parallel_size: int = 1,
sequence_model_parallel_size: int = 1,
data_parallel_size: int = 1,
backend: Optional[str] = None,
) -> None:
"""
@@ -813,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)
@@ -862,14 +859,13 @@ def initialize_model_parallel(
group_name="sp")
# Build the data parallel groups.
num_data_parallel_groups: int = (world_size // data_parallel_size)
num_data_parallel_groups: int = sequence_model_parallel_size
global _DP
assert _DP is None, ("data parallel group is already initialized")
group_ranks = []
for i in range(num_data_parallel_groups):
ranks = list(range(i * data_parallel_size,
(i + 1) * data_parallel_size))
ranks = list(range(i, world_size, num_data_parallel_groups))
group_ranks.append(ranks)
_DP = init_model_parallel_group(group_ranks,
@@ -878,56 +874,62 @@ def initialize_model_parallel(
group_name="dp")
def get_sequence_model_parallel_world_size() -> int:
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 get_data_parallel_world_size() -> int:
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_data_parallel_rank() -> int:
def get_dp_rank() -> int:
"""Return my rank for the data parallel group."""
return get_dp_group().rank_in_group
def ensure_model_parallel_initialized(
tensor_model_parallel_size: int,
sequence_model_parallel_size: int,
data_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,
data_parallel_size, backend)
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:
@@ -963,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
+1 -5
View File
@@ -94,11 +94,7 @@ class VideoGenerator:
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()
fastvideo_args = FastVideoArgs(model_path=model_path, **config_args)
return cls.from_fastvideo_args(fastvideo_args)
+38 -48
View File
@@ -42,10 +42,10 @@ class FastVideoArgs:
# Parallelism
num_gpus: int = 1
tp_size: Optional[int] = None
sp_size: Optional[int] = None
dp_size: int = 1
dp_shards: 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
@@ -85,8 +85,8 @@ class FastVideoArgs:
postprocess_text_funcs: Tuple[Callable[[Any], Any], ...] = field(
default_factory=lambda: (postprocess_text, ))
# STA (Spatial-Temporal Attention) parameters
STA_mode: str = "STA_inference"
# STA parameters
STA_mode: Optional[str] = None
skip_time_steps: int = 15
# LoRA parameters
lora_path: Optional[str] = None
@@ -109,16 +109,12 @@ class FastVideoArgs:
# Logging
log_level: str = "info"
# Inference parameters
device_str: Optional[str] = None
device = None
@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:
@@ -193,17 +189,15 @@ class FastVideoArgs:
help="The sequence parallelism size.",
)
parser.add_argument(
"--data-parallel-size",
"--dp-size",
"--hsdp-replicate-dim",
type=int,
default=FastVideoArgs.dp_size,
default=FastVideoArgs.hsdp_replicate_dim,
help="The data parallelism size.",
)
parser.add_argument(
"--data-parallel-shards",
"--dp-shards",
"--hsdp-shard-dim",
type=int,
default=FastVideoArgs.dp_shards,
default=FastVideoArgs.hsdp_shard_dim,
help="The data parallelism shards.",
)
parser.add_argument(
@@ -280,13 +274,14 @@ class FastVideoArgs:
help="Precision for image encoder",
)
# STA (Spatial-Temporal Attention) parameters
# STA parameters
parser.add_argument(
"--STA-mode",
type=str,
default=FastVideoArgs.STA_mode,
choices=[
"STA_inference", "STA_searching", "STA_tuning", "STA_tuning_cfg"
"STA_inference", "STA_searching", "STA_tuning",
"STA_tuning_cfg", None
],
help="STA mode",
)
@@ -382,35 +377,34 @@ class FastVideoArgs:
kwargs[attr] = args.tensor_parallel_size
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
kwargs[attr] = args.sequence_parallel_size
elif attr == 'dp_size' and hasattr(args, 'data_parallel_size'):
kwargs[attr] = args.data_parallel_size
elif attr == 'dp_shards' and hasattr(args, 'data_parallel_shards'):
kwargs[attr] = args.data_parallel_shards
elif attr == 'flow_shift' and hasattr(args, 'shift'):
kwargs[attr] = args.shift
# 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)
if value is not None:
kwargs[attr] = value
return cls(**kwargs)
def check_fastvideo_args(self) -> None:
"""Validate inference arguments for consistency"""
if not self.inference_mode:
assert self.dp_size is not None, "dp_size must be set for training"
assert self.dp_shards is not None, "dp_shards must be set for training"
assert self.sp_size is not None, "sp_size must be set for training"
assert self.hsdp_replicate_dim != -1, "hsdp_replicate_dim must be set for training"
assert self.hsdp_shard_dim != -1, "hsdp_shard_dim must be set for training"
assert self.sp_size != -1, "sp_size must be set for training"
if self.tp_size is None:
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.dp_shards is None:
self.dp_shards = self.num_gpus
if self.hsdp_shard_dim == -1:
self.hsdp_shard_dim = self.num_gpus
assert self.sp_size <= self.num_gpus and self.num_gpus % self.sp_size == 0, "num_gpus must >= and be divisible by sp_size"
assert self.dp_size <= self.num_gpus and self.num_gpus % self.dp_size == 0, "num_gpus must >= and be divisible by dp_size"
assert self.dp_shards <= self.num_gpus and self.num_gpus % self.dp_shards == 0, "num_gpus must >= and be divisible by dp_shards"
assert self.hsdp_replicate_dim <= self.num_gpus and self.num_gpus % self.hsdp_replicate_dim == 0, "num_gpus must >= and be divisible by hsdp_replicate_dim"
assert self.hsdp_shard_dim <= self.num_gpus and self.num_gpus % self.hsdp_shard_dim == 0, "num_gpus must >= and be divisible by hsdp_shard_dim"
if self.num_gpus < max(self.tp_size, self.sp_size):
self.num_gpus = max(self.tp_size, self.sp_size)
@@ -466,7 +460,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
@@ -530,7 +523,8 @@ class TrainingArgs(FastVideoArgs):
precondition_outputs: bool = False
# validation & logs
validation_prompt_dir: str = ""
validation_dataset_file: str = ""
validation_path: str = ""
validation_sampling_steps: str = ""
validation_guidance_scale: str = ""
validation_steps: float = 0.0
@@ -543,7 +537,6 @@ class TrainingArgs(FastVideoArgs):
checkpoints_total_limit: int = 0
checkpointing_steps: int = 0
resume_from_checkpoint: bool = False
logging_dir: str = ""
# optimizer & scheduler
num_train_epochs: int = 0
@@ -551,7 +544,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
@@ -584,9 +577,6 @@ class TrainingArgs(FastVideoArgs):
# master_weight_type
master_weight_type: str = ""
# For fast checking in LoRA pipeline
training_mode: bool = True
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
# Get all fields from the dataclass
@@ -602,14 +592,10 @@ class TrainingArgs(FastVideoArgs):
kwargs[attr] = args.sequence_parallel_size
elif attr == 'flow_shift' and hasattr(args, 'shift'):
kwargs[attr] = args.shift
elif attr == 'dp_size' and hasattr(args, 'data_parallel_size'):
kwargs[attr] = args.data_parallel_size
elif attr == 'dp_shards' and hasattr(args, 'data_parallel_shards'):
kwargs[attr] = args.data_parallel_shards
# 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)
if getattr(args, attr, default_value) is not None:
kwargs[attr] = getattr(args, attr, default_value)
return cls(**kwargs)
@@ -683,9 +669,12 @@ class TrainingArgs(FastVideoArgs):
help="Whether to precondition the outputs of the model")
# Validation and logging
parser.add_argument("--validation-prompt-dir",
parser.add_argument("--validation-dataset-file",
type=str,
help="Directory containing validation prompts")
help="File containing validation dataset")
parser.add_argument("--validation-path",
type=str,
help="Path to validation dataset")
parser.add_argument("--validation-sampling-steps",
type=str,
help="Validation sampling steps")
@@ -703,6 +692,7 @@ class TrainingArgs(FastVideoArgs):
help="Project name for tracking")
parser.add_argument("--seed",
type=int,
default=42,
help="Seed for deterministic training")
# Output configuration
+4 -4
View File
@@ -5,17 +5,16 @@ import time
from collections import defaultdict
from contextlib import contextmanager
from dataclasses import dataclass
from typing import TYPE_CHECKING, Optional
from typing import Optional
import torch
# if TYPE_CHECKING:
from fastvideo.v1.attention import AttentionMetadata
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
if TYPE_CHECKING:
from fastvideo.v1.attention import AttentionMetadata
logger = init_logger(__name__)
# TODO(will): check if this is needed
@@ -70,6 +69,7 @@ def set_forward_context(current_timestep,
_forward_context = ForwardContext(current_timestep=current_timestep,
attn_metadata=attn_metadata,
forward_batch=forward_batch)
try:
yield
finally:
+14 -15
View File
@@ -8,8 +8,7 @@ import torch
import torch.nn.functional as F
from torch.nn.parameter import Parameter
from fastvideo.v1.distributed import (divide, get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
from fastvideo.v1.distributed import (divide, get_tp_rank, get_tp_world_size,
split_tensor_along_last_dim,
tensor_model_parallel_all_gather,
tensor_model_parallel_all_reduce)
@@ -273,7 +272,7 @@ class ColumnParallelLinear(LinearBase):
output_sizes: Optional[list[int]] = None,
prefix: str = ""):
# Divide the weight matrix along the last dimension.
self.tp_size = get_tensor_model_parallel_world_size()
self.tp_size = get_tp_world_size()
self.input_size_per_partition = input_size
self.output_size_per_partition = divide(output_size, self.tp_size)
self.output_partition_sizes = [self.output_size_per_partition]
@@ -315,7 +314,7 @@ class ColumnParallelLinear(LinearBase):
def weight_loader(self, param: Parameter,
loaded_weight: torch.Tensor) -> None:
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_tp_rank()
output_dim = getattr(param, "output_dim", None)
is_sharded_weight = getattr(param, "is_sharded_weight", False)
@@ -365,7 +364,7 @@ class ColumnParallelLinear(LinearBase):
s = f"in_features={self.input_size}"
s += f", output_features={self.output_size_per_partition}"
s += f", bias={self.bias is not None}"
s += f", tp_size={get_tensor_model_parallel_world_size()}"
s += f", tp_size={get_tp_world_size()}"
s += f", gather_output={self.gather_output}"
return s
@@ -403,7 +402,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
self.output_sizes = output_sizes
tp_size = get_tensor_model_parallel_world_size()
tp_size = get_tp_world_size()
assert all(output_size % tp_size == 0 for output_size in output_sizes)
super().__init__(input_size=input_size,
output_size=sum(output_sizes),
@@ -449,8 +448,8 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
return
assert loaded_shard_id < len(self.output_sizes)
tp_rank = get_tensor_model_parallel_rank()
tp_size = get_tensor_model_parallel_world_size()
tp_rank = get_tp_rank()
tp_size = get_tp_world_size()
if output_dim is not None:
shard_offset = sum(self.output_sizes[:loaded_shard_id]) // tp_size
shard_size = self.output_sizes[loaded_shard_id] // tp_size
@@ -540,7 +539,7 @@ class MergedColumnParallelLinear(ColumnParallelLinear):
assert loaded_shard_id < len(self.output_sizes)
tp_size = get_tensor_model_parallel_world_size()
tp_size = get_tp_world_size()
if isinstance(param, BlockQuantScaleParameter):
raise NotImplementedError("FP8 is not implemented yet")
@@ -611,7 +610,7 @@ class QKVParallelLinear(ColumnParallelLinear):
total_num_kv_heads = total_num_heads
self.total_num_kv_heads = total_num_kv_heads
# Divide the weight matrix along the last dimension.
tp_size = get_tensor_model_parallel_world_size()
tp_size = get_tp_world_size()
self.num_heads = divide(self.total_num_heads, tp_size)
if tp_size >= self.total_num_kv_heads:
self.num_kv_heads = 1
@@ -757,7 +756,7 @@ class QKVParallelLinear(ColumnParallelLinear):
self.weight_loader(param, loaded_weight_shard, shard_id)
return
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_tp_rank()
assert loaded_shard_id in ["q", "k", "v"]
# If output dim is defined, use the default loading process.
@@ -850,8 +849,8 @@ class RowParallelLinear(LinearBase):
quant_config: Optional[QuantizationConfig] = None,
prefix: str = ""):
# Divide the weight matrix along the first dimension.
self.tp_rank = get_tensor_model_parallel_rank()
self.tp_size = get_tensor_model_parallel_world_size()
self.tp_rank = get_tp_rank()
self.tp_size = get_tp_world_size()
self.input_size_per_partition = divide(input_size, self.tp_size)
self.output_size_per_partition = output_size
self.output_partition_sizes = [output_size]
@@ -888,7 +887,7 @@ class RowParallelLinear(LinearBase):
self.register_parameter("bias", None)
def weight_loader(self, param: Parameter, loaded_weight: torch.Tensor):
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_tp_rank()
input_dim = getattr(param, "input_dim", None)
is_sharded_weight = getattr(param, "is_sharded_weight", False)
# bitsandbytes loads the weights of the specific portion
@@ -925,7 +924,7 @@ class RowParallelLinear(LinearBase):
if self.input_is_parallel:
input_parallel = input_
else:
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_tp_rank()
splitted_input = split_tensor_along_last_dim(
input_, num_partitions=self.tp_size)
input_parallel = splitted_input[tp_rank].contiguous()
+7 -7
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
# Code adapted from SGLang https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/lora/layers.py
from typing import Dict, List, Tuple, Type, Union
@@ -6,8 +7,7 @@ import torch
from torch import nn
from torch.distributed.tensor import DTensor, distribute_tensor
from fastvideo.v1.distributed import (get_tensor_model_parallel_rank,
split_tensor_along_last_dim,
from fastvideo.v1.distributed import (get_tp_rank, split_tensor_along_last_dim,
tensor_model_parallel_all_gather,
tensor_model_parallel_all_reduce)
from fastvideo.v1.layers.linear import (ColumnParallelLinear, LinearBase,
@@ -160,7 +160,7 @@ class ColumnParallelLinearWithLoRA(BaseLayerWithLoRA):
return A
def slice_lora_b_weights(self, B: torch.Tensor) -> torch.Tensor:
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_tp_rank()
shard_size = self.base_layer.output_partition_sizes[0]
start_idx = tp_rank * shard_size
end_idx = (tp_rank + 1) * shard_size
@@ -180,7 +180,7 @@ class MergedColumnParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
return A.to(self.base_layer.weight)
def slice_lora_b_weights(self, B: torch.Tensor) -> torch.Tensor:
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_tp_rank()
# Since the outputs for both gate and up are identical, we use a random one.
shard_size = self.base_layer.output_partition_sizes[0]
start_idx = tp_rank * shard_size
@@ -201,7 +201,7 @@ class QKVParallelLinearWithLoRA(ColumnParallelLinearWithLoRA):
def slice_lora_b_weights(
self, B: List[torch.Tensor]) -> Tuple[torch.Tensor, torch.Tensor]:
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_tp_rank()
B_q, B_kv = B
base_layer = self.base_layer
q_proj_shard_size = base_layer.q_proj_shard_size
@@ -232,7 +232,7 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
if self.base_layer.input_is_parallel:
input_parallel = input_
else:
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_tp_rank()
splitted_input = split_tensor_along_last_dim(
input_, num_partitions=self.base_layer.tp_size)
input_parallel = splitted_input[tp_rank].contiguous()
@@ -257,7 +257,7 @@ class RowParallelLinearWithLoRA(BaseLayerWithLoRA):
return output, output_bias
def slice_lora_a_weights(self, A: torch.Tensor) -> torch.Tensor:
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_tp_rank()
shard_size = self.base_layer.input_size_per_partition
start_idx = tp_rank * shard_size
end_idx = (tp_rank + 1) * shard_size
@@ -7,8 +7,7 @@ import torch
import torch.nn.functional as F
from torch.nn.parameter import Parameter, UninitializedParameter
from fastvideo.v1.distributed import (divide, get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
from fastvideo.v1.distributed import (divide, get_tp_rank, get_tp_world_size,
tensor_model_parallel_all_reduce)
from fastvideo.v1.layers.quantization.base_config import (
QuantizationConfig, QuantizeMethodBase, method_has_implemented_embedding)
@@ -205,8 +204,8 @@ class VocabParallelEmbedding(torch.nn.Module):
super().__init__()
# Keep the input dimensions.
tp_rank = get_tensor_model_parallel_rank()
self.tp_size = get_tensor_model_parallel_world_size()
tp_rank = get_tp_rank()
self.tp_size = get_tp_world_size()
self.num_embeddings = num_embeddings
self.padding_size = padding_size
self.org_vocab_size = org_num_embeddings or num_embeddings
+9 -13
View File
@@ -9,8 +9,7 @@ import torch.nn as nn
from fastvideo.v1.attention import DistributedAttention, LocalAttention
from fastvideo.v1.configs.models.dits import HunyuanVideoConfig
from fastvideo.v1.configs.sample.teacache import TeaCacheParams
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.distributed.parallel_state import get_sp_world_size
from fastvideo.v1.forward_context import get_forward_context
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
ScaleResidualLayerNormScaleShift)
@@ -591,9 +590,8 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
# Get rotary embeddings
freqs_cos, freqs_sin = get_rotary_pos_embed(
(tt * get_sequence_model_parallel_world_size(), th, tw),
self.hidden_size, self.num_attention_heads, self.rope_dim_list,
self.rope_theta)
(tt * get_sp_world_size(), th, tw), self.hidden_size,
self.num_attention_heads, self.rope_dim_list, self.rope_theta)
freqs_cos = freqs_cos.to(x.device)
freqs_sin = freqs_sin.to(x.device)
# Prepare modulation vectors
@@ -689,18 +687,16 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
# convert to DTensor
vec_ = torch.distributed.tensor.DTensor.from_local(
vec_,
torch.distributed.DeviceMesh(
"cuda",
list(range(get_sequence_model_parallel_world_size())),
mesh_dim_names=("dp", )),
torch.distributed.DeviceMesh("cuda",
list(range(get_sp_world_size())),
mesh_dim_names=("dp", )),
[torch.distributed.tensor.Replicate()])
inp = torch.distributed.tensor.DTensor.from_local(
inp,
torch.distributed.DeviceMesh(
"cuda",
list(range(get_sequence_model_parallel_world_size())),
mesh_dim_names=("dp", )),
torch.distributed.DeviceMesh("cuda",
list(range(get_sp_world_size())),
mesh_dim_names=("dp", )),
[torch.distributed.tensor.Replicate()])
# txt_ = kwargs["txt"].clone()
+2 -4
View File
@@ -18,8 +18,7 @@ from torch import nn
from fastvideo.v1.attention import DistributedAttention, LocalAttention
from fastvideo.v1.configs.models.dits import StepVideoConfig
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.distributed.parallel_state import get_sp_world_size
from fastvideo.v1.layers.layernorm import LayerNormScaleShift
from fastvideo.v1.layers.linear import ReplicatedLinear
from fastvideo.v1.layers.mlp import MLP
@@ -578,8 +577,7 @@ class StepVideoModel(BaseDiT):
key = (F, Ht, W, dtype)
if key not in self._rope_cache:
cos, sin = get_rotary_pos_embed(
rope_sizes=(F * get_sequence_model_parallel_world_size(), Ht,
W),
rope_sizes=(F * get_sp_world_size(), Ht, W),
hidden_size=self.hidden_size,
heads_num=self.hidden_size // self.attention_head_dim,
rope_dim_list=(64, 32, 32), # same split you used
+169 -14
View File
@@ -7,11 +7,12 @@ import numpy as np
import torch
import torch.nn as nn
from fastvideo.v1.attention import DistributedAttention, LocalAttention
import fastvideo.v1.envs as envs
from fastvideo.v1.attention import (DistributedAttention,
DistributedAttention_VSA, LocalAttention)
from fastvideo.v1.configs.models.dits import WanVideoConfig
from fastvideo.v1.configs.sample.wan import WanTeaCacheParams
from fastvideo.v1.distributed.parallel_state import (
get_sequence_model_parallel_world_size)
from fastvideo.v1.distributed.parallel_state import get_sp_world_size
from fastvideo.v1.forward_context import get_forward_context
from fastvideo.v1.layers.layernorm import (LayerNormScaleShift, RMSNorm,
ScaleResidual,
@@ -233,6 +234,7 @@ class WanTransformerBlock(nn.Module):
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
self.to_out = ReplicatedLinear(dim, dim, bias=True)
self.attn1 = DistributedAttention(
num_heads=num_heads,
@@ -354,6 +356,155 @@ class WanTransformerBlock(nn.Module):
return hidden_states
class WanTransformerBlock_VSA(nn.Module):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
prefix: str = ""):
super().__init__()
# 1. Self-attention
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
self.to_gate_compress = ReplicatedLinear(dim, dim, bias=True)
self.to_out = ReplicatedLinear(dim, dim, bias=True)
self.attn1 = DistributedAttention_VSA(
num_heads=num_heads,
head_size=dim // num_heads,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn1")
self.hidden_dim = dim
self.num_attention_heads = num_heads
dim_head = dim // num_heads
if qk_norm == "rms_norm":
self.norm_q = RMSNorm(dim_head, eps=eps)
self.norm_k = RMSNorm(dim_head, eps=eps)
elif qk_norm == "rms_norm_across_heads":
# LTX applies qk norm across all heads
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
else:
print("QK Norm type not supported")
raise Exception
assert cross_attn_norm is True
self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
# I2V
self.attn2 = WanI2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
else:
# T2V
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
self.mlp_residual = ScaleResidual()
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,
freqs_cis: Tuple[torch.Tensor, torch.Tensor],
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
gate_compress, _ = self.to_gate_compress(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
gate_compress = gate_compress.squeeze(1).unflatten(
2, (self.num_attention_heads, -1))
# Apply rotary embeddings
cos, sin = freqs_cis
query, key = _apply_rotary_emb(query, cos, sin,
is_neox_style=False), _apply_rotary_emb(
key, cos, sin, is_neox_style=False)
attn_output, _ = self.attn1(query,
key,
value,
gate_compress=gate_compress)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
context=encoder_hidden_states,
context_lens=None)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
class WanTransformer3DModel(CachableDiT):
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
_compile_conditions = WanVideoConfig()._compile_conditions
@@ -390,16 +541,18 @@ class WanTransformer3DModel(CachableDiT):
)
# 3. Transformer blocks
attn_backend = envs.FASTVIDEO_ATTENTION_BACKEND
transformer_block = WanTransformerBlock_VSA if attn_backend == "VIDEO_SPARSE_ATTN" else WanTransformerBlock
self.blocks = nn.ModuleList([
WanTransformerBlock(inner_dim,
config.ffn_dim,
config.num_attention_heads,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
self._supported_attention_backends,
prefix=f"{config.prefix}.blocks.{i}")
transformer_block(inner_dim,
config.ffn_dim,
config.num_attention_heads,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
self._supported_attention_backends,
prefix=f"{config.prefix}.blocks.{i}")
for i in range(config.num_layers)
])
@@ -448,8 +601,8 @@ class WanTransformer3DModel(CachableDiT):
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sequence_model_parallel_world_size(),
post_patch_height, post_patch_width),
(post_patch_num_frames * get_sp_world_size(), post_patch_height,
post_patch_width),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
@@ -465,6 +618,8 @@ class WanTransformer3DModel(CachableDiT):
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image)
if encoder_hidden_states.dim() == 2:
encoder_hidden_states = encoder_hidden_states.unsqueeze(0)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
from abc import ABC, abstractmethod
from typing import Optional, Tuple
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
# type: ignore
import os
+2 -3
View File
@@ -13,8 +13,7 @@ from fastvideo.v1.attention import LocalAttention
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
CLIPTextConfig,
CLIPVisionConfig)
from fastvideo.v1.distributed import (divide,
get_tensor_model_parallel_world_size)
from fastvideo.v1.distributed import divide, get_tp_world_size
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.layers.linear import (ColumnParallelLinear, QKVParallelLinear,
RowParallelLinear)
@@ -160,7 +159,7 @@ class CLIPAttention(nn.Module):
prefix=f"{prefix}.out_proj",
)
self.tp_size = get_tensor_model_parallel_world_size()
self.tp_size = get_tp_world_size()
self.num_heads_per_partition = divide(self.num_heads, self.tp_size)
self.attn = LocalAttention(
+2 -2
View File
@@ -32,7 +32,7 @@ from torch import nn
from fastvideo.v1.attention import LocalAttention
# from ..utils import (extract_layer_index)
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput, LlamaConfig
from fastvideo.v1.distributed import get_tensor_model_parallel_world_size
from fastvideo.v1.distributed import get_tp_world_size
from fastvideo.v1.layers.activation import SiluAndMul
from fastvideo.v1.layers.layernorm import RMSNorm
from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
@@ -101,7 +101,7 @@ class LlamaAttention(nn.Module):
super().__init__()
# layer_idx = extract_layer_index(prefix)
self.hidden_size = hidden_size
tp_size = get_tensor_model_parallel_world_size()
tp_size = get_tp_world_size()
self.total_num_heads = num_heads
assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
# type: ignore
# Copyright 2025 StepFun Inc. All Rights Reserved.
#
+4 -5
View File
@@ -28,8 +28,7 @@ import torch.nn.functional as F
from torch import nn
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput, T5Config
from fastvideo.v1.distributed import (get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size)
from fastvideo.v1.distributed import get_tp_rank, get_tp_world_size
from fastvideo.v1.layers.activation import get_act_fn
from fastvideo.v1.layers.layernorm import RMSNorm
from fastvideo.v1.layers.linear import (MergedColumnParallelLinear,
@@ -171,7 +170,7 @@ class T5Attention(nn.Module):
self.total_num_heads = self.total_num_kv_heads = config.num_heads
# Partition heads across multiple tensor parallel GPUs.
tp_world_size = get_tensor_model_parallel_world_size()
tp_world_size = get_tp_world_size()
assert config.num_heads % tp_world_size == 0
self.n_heads = config.num_heads // tp_world_size
@@ -329,8 +328,8 @@ class T5Attention(nn.Module):
attn_bias.masked_fill_(attention_mask == 0,
torch.finfo(q.dtype).min)
if get_tensor_model_parallel_world_size() > 1:
rank = get_tensor_model_parallel_rank()
if get_tp_world_size() > 1:
rank = get_tp_rank()
attn_bias = attn_bias[:, rank * self.n_heads:(rank + 1) *
self.n_heads, :, :]
attn_output = self.attn(q, k, v, attn_bias)
@@ -15,6 +15,7 @@ from safetensors.torch import load_file as safetensors_load_file
from transformers import AutoImageProcessor, AutoTokenizer
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.hf_transformer_utils import get_diffusers_config
@@ -227,7 +228,7 @@ class TextEncoderLoader(ComponentLoader):
encoder_config.update_model_arch(model_config)
encoder_precision = fastvideo_args.text_encoder_precisions[1]
target_device = torch.device(fastvideo_args.device_str)
target_device = get_torch_device()
# TODO(will): add support for other dtypes
return self.load_model(model_path, encoder_config, target_device,
encoder_precision)
@@ -286,7 +287,7 @@ class ImageEncoderLoader(TextEncoderLoader):
encoder_config = fastvideo_args.image_encoder_config
encoder_config.update_model_arch(model_config)
target_device = torch.device(fastvideo_args.device_str)
target_device = get_torch_device()
# TODO(will): add support for other dtypes
return self.load_model(model_path, encoder_config, target_device,
fastvideo_args.image_encoder_precision)
@@ -342,7 +343,7 @@ class VAELoader(ComponentLoader):
vae_config.update_model_arch(config)
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(fastvideo_args.device)
vae = vae_cls(vae_config).to(get_torch_device())
# Find all safetensors files
safetensors_list = glob.glob(
@@ -402,7 +403,7 @@ class TransformerLoader(ComponentLoader):
# Load the model using FSDP loader
logger.info("Loading model from %s, default_dtype: %s", cls_name,
default_dtype)
assert fastvideo_args.dp_shards is not None
assert fastvideo_args.hsdp_shard_dim is not None
model = maybe_load_fsdp_model(
model_cls=model_cls,
init_params={
@@ -410,9 +411,9 @@ class TransformerLoader(ComponentLoader):
"hf_config": hf_config
},
weight_dir_list=safetensors_list,
device=fastvideo_args.device,
data_parallel_size=fastvideo_args.dp_size,
data_parallel_shards=fastvideo_args.dp_shards,
device=get_torch_device(),
hsdp_replicate_dim=fastvideo_args.hsdp_replicate_dim,
hsdp_shard_dim=fastvideo_args.hsdp_shard_dim,
cpu_offload=fastvideo_args.use_cpu_offload,
fsdp_inference=fastvideo_args.use_fsdp_inference,
default_dtype=default_dtype,
+48 -8
View File
@@ -60,8 +60,8 @@ def maybe_load_fsdp_model(
init_params: Dict[str, Any],
weight_dir_list: List[str],
device: torch.device,
data_parallel_size: int,
data_parallel_shards: int,
hsdp_replicate_dim: int,
hsdp_shard_dim: int,
default_dtype: torch.dtype,
param_dtype: torch.dtype,
reduce_dtype: torch.dtype,
@@ -87,13 +87,15 @@ def maybe_load_fsdp_model(
with set_default_dtype(default_dtype), torch.device("meta"):
model = model_cls(**init_params)
dp_size = data_parallel_size if fsdp_inference or training_mode else 1
world_size = hsdp_replicate_dim * hsdp_shard_dim
if not training_mode and not fsdp_inference:
hsdp_replicate_dim = world_size
hsdp_shard_dim = 1
device_mesh = init_device_mesh(
"cuda",
# (Replicate(), Shard(dim=0))
mesh_shape=(dp_size, data_parallel_shards),
mesh_dim_names=("dp", "sp"),
mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim),
mesh_dim_names=("replicate", "shard"),
)
shard_model(model,
cpu_offload=cpu_offload,
@@ -216,13 +218,15 @@ def load_model_from_full_model_state_dict(
NotImplementedError: If got FSDP with more than 1D.
"""
meta_sd = model.state_dict()
# Find new params
used_keys = set()
sharded_sd = {}
to_merge_params: DefaultDict[str, Dict[Any, Any]] = defaultdict(dict)
for source_param_name, full_tensor in full_sd_iterator:
assert param_names_mapping is not None
target_param_name, merge_index, num_params_to_merge = param_names_mapping(
source_param_name)
used_keys.add(target_param_name)
if merge_index is not None:
to_merge_params[target_param_name][merge_index] = full_tensor
if len(to_merge_params[target_param_name]) == num_params_to_merge:
@@ -241,7 +245,6 @@ def load_model_from_full_model_state_dict(
raise ValueError(
f"Parameter {source_param_name}-->{target_param_name} not found in meta sharded state dict"
)
if not hasattr(meta_sharded_param, "device_mesh"):
full_tensor = full_tensor.to(device=device, dtype=param_dtype)
# In cases where parts of the model aren't sharded, some parameters will be plain tensors
@@ -256,5 +259,42 @@ def load_model_from_full_model_state_dict(
if cpu_offload:
sharded_tensor = sharded_tensor.cpu()
sharded_sd[target_param_name] = nn.Parameter(sharded_tensor)
unused_keys = set(meta_sd.keys()) - used_keys
if unused_keys:
logger.warning("Found new parameters in meta state dict: %s",
unused_keys)
# List of allowed parameter name patterns
ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress"] # Can be extended as needed
for new_param_name in unused_keys:
if not any(pattern in new_param_name
for pattern in ALLOWED_NEW_PARAM_PATTERNS):
logger.error("Unsupported new parameter: %s. Allowed patterns: %s",
new_param_name, ALLOWED_NEW_PARAM_PATTERNS)
raise ValueError(
f"New parameter '{new_param_name}' is not supported. "
f"Currently only parameters containing {ALLOWED_NEW_PARAM_PATTERNS} are allowed."
)
meta_sharded_param = meta_sd.get(new_param_name)
if not hasattr(meta_sharded_param, "device_mesh"):
# Initialize with zeros
sharded_tensor = torch.zeros_like(meta_sharded_param,
device=device,
dtype=param_dtype)
else:
# Initialize with zeros and distribute
full_tensor = torch.zeros_like(meta_sharded_param,
device=device,
dtype=param_dtype)
sharded_tensor = distribute_tensor(
full_tensor,
meta_sharded_param.device_mesh,
meta_sharded_param.placements,
)
if cpu_offload:
sharded_tensor = sharded_tensor.cpu()
sharded_sd[new_param_name] = nn.Parameter(sharded_tensor)
# choose `assign=True` since we cannot call `copy_` on meta tensor
return model.load_state_dict(sharded_sd, strict=strict, assign=True)
-107
View File
@@ -1,20 +1,16 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/model_executor/model_loader/weight_utils.py
"""Utilities for downloading and initializing model weights."""
import fnmatch
import hashlib
import json
import os
import tempfile
import time
from collections import defaultdict
from pathlib import Path
from typing import Generator, List, Optional, Tuple, Union
import filelock
import huggingface_hub.constants
import torch
from huggingface_hub import HfFileSystem, hf_hub_download, snapshot_download
from safetensors.torch import safe_open
from tqdm.auto import tqdm
@@ -64,109 +60,6 @@ def get_lock(model_name_or_path: Union[str, Path],
return lock
def _shared_pointers(tensors):
ptrs = defaultdict(list)
for k, v in tensors.items():
ptrs[v.data_ptr()].append(k)
failing = []
for _, names in ptrs.items():
if len(names) > 1:
failing.append(names)
return failing
def download_weights_from_hf(
model_name_or_path: str,
cache_dir: Optional[str],
allow_patterns: List[str],
revision: Optional[str] = None,
ignore_patterns: Optional[Union[str, List[str]]] = None,
) -> str:
"""Download model weights from Hugging Face Hub.
Args:
model_name_or_path (str): The model name or path.
cache_dir (Optional[str]): The cache directory to store the model
weights. If None, will use HF defaults.
allow_patterns (List[str]): The allowed patterns for the
weight files. Files matched by any of the patterns will be
downloaded.
revision (Optional[str]): The revision of the model.
ignore_patterns (Optional[Union[str, List[str]]]): The patterns to
filter out the weight files. Files matched by any of the patterns
will be ignored.
Returns:
str: The path to the downloaded model weights.
"""
local_only = huggingface_hub.constants.HF_HUB_OFFLINE
if not local_only:
# Before we download we look at that is available:
fs = HfFileSystem()
file_list = fs.ls(model_name_or_path, detail=False, revision=revision)
# depending on what is available we download different things
for pattern in allow_patterns:
matching = fnmatch.filter(file_list, pattern)
if len(matching) > 0:
allow_patterns = [pattern]
break
logger.info("Using model weights format %s", allow_patterns)
# Use file lock to prevent multiple processes from
# downloading the same model weights at the same time.
with get_lock(model_name_or_path, cache_dir):
start_time = time.perf_counter()
hf_folder: str = snapshot_download(
model_name_or_path,
allow_patterns=allow_patterns,
ignore_patterns=ignore_patterns,
cache_dir=cache_dir,
tqdm_class=DisabledTqdm,
revision=revision,
local_files_only=local_only,
)
time_taken = time.perf_counter() - start_time
if time_taken > 0.5:
logger.info("Time spent downloading weights for %s: %.6f seconds",
model_name_or_path, time_taken)
return hf_folder
def download_safetensors_index_file_from_hf(
model_name_or_path: str,
index_file: str,
cache_dir: Optional[str],
revision: Optional[str] = None,
) -> None:
"""Download hf safetensors index file from Hugging Face Hub.
Args:
model_name_or_path (str): The model name or path.
cache_dir (Optional[str]): The cache directory to store the model
weights. If None, will use HF defaults.
revision (Optional[str]): The revision of the model.
"""
# Use file lock to prevent multiple processes from
# downloading the same model weights at the same time.
with get_lock(model_name_or_path, cache_dir):
try:
# Download the safetensors index file.
hf_hub_download(
repo_id=model_name_or_path,
filename=index_file,
cache_dir=cache_dir,
revision=revision,
local_files_only=huggingface_hub.constants.HF_HUB_OFFLINE,
)
# If file not found on remote or locally, we should not fail since
# only some models will have index_file.
except huggingface_hub.utils.EntryNotFoundError:
logger.info("No %s found in remote.", index_file)
except huggingface_hub.utils.LocalEntryNotFoundError:
logger.info("No %s found in local cache.", index_file)
# For models like Mistral-7B-v0.3, there are both sharded
# safetensors files and a consolidated safetensors file.
# Passing both of these to the weight loader functionality breaks.
+5 -5
View File
@@ -7,7 +7,7 @@ from typing import Any, Callable, Tuple, Union
import torch
from torch.nn import Parameter
from fastvideo.v1.distributed import get_tensor_model_parallel_rank
from fastvideo.v1.distributed import get_tp_rank
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.utils import _make_synced_weight_loader
@@ -97,7 +97,7 @@ class _ColumnvLLMParameter(BasevLLMParameter):
return self._output_dim
def load_column_parallel_weight(self, loaded_weight: torch.Tensor) -> None:
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_tp_rank()
shard_size = self.data.shape[self.output_dim]
loaded_weight = loaded_weight.narrow(self.output_dim,
tp_rank * shard_size, shard_size)
@@ -120,7 +120,7 @@ class _ColumnvLLMParameter(BasevLLMParameter):
param_data = self.data
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_tp_rank()
param_data = param_data.narrow(self.output_dim, shard_offset,
shard_size)
loaded_weight = loaded_weight.narrow(self.output_dim,
@@ -148,7 +148,7 @@ class _ColumnvLLMParameter(BasevLLMParameter):
shard_offset=shard_offset, shard_size=shard_size)
param_data = self.data
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_tp_rank()
shard_id = tp_rank if shard_id == "q" else tp_rank // num_heads
param_data = param_data.narrow(self.output_dim, shard_offset,
shard_size)
@@ -176,7 +176,7 @@ class RowvLLMParameter(BasevLLMParameter):
return self._input_dim
def load_row_parallel_weight(self, loaded_weight: torch.Tensor) -> None:
tp_rank = get_tensor_model_parallel_rank()
tp_rank = get_tp_rank()
shard_size = self.data.shape[self.input_dim]
loaded_weight = loaded_weight.narrow(self.input_dim,
tp_rank * shard_size, shard_size)
@@ -25,11 +25,12 @@ from typing import Any, Optional, Tuple, Union
import torch
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput, logging
from diffusers.utils import BaseOutput
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.schedulers.base import BaseScheduler
logger = logging.get_logger(__name__) # pylint: disable=invalid-name
logger = init_logger(__name__)
@dataclass
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
# Copied from https://github.com/huggingface/diffusers/blob/v0.31.0/src/diffusers/schedulers/scheduling_unipc_multistep.py
# Convert unipc for flow matching
# Copyright 2024-2025 The Alibaba Wan Team Authors. All rights reserved.
+3 -5
View File
@@ -10,8 +10,7 @@ import torch.distributed as dist
from diffusers.utils.torch_utils import randn_tensor
from fastvideo.v1.configs.models import VAEConfig
from fastvideo.v1.distributed import (get_sequence_model_parallel_rank,
get_sequence_model_parallel_world_size)
from fastvideo.v1.distributed import get_sp_parallel_rank, get_sp_world_size
class ParallelTiledVAE(ABC):
@@ -84,7 +83,7 @@ class ParallelTiledVAE(ABC):
num_sample_frames = (num_frames -
1) * self.temporal_compression_ratio + 1
if self.use_tiling and self.use_parallel_tiling and get_sequence_model_parallel_world_size(
if self.use_tiling and self.use_parallel_tiling and get_sp_world_size(
) > 1:
return self.parallel_tiled_decode(z)[:, :, :num_sample_frames]
if self.use_tiling and self.use_temporal_tiling and num_frames > tile_latent_min_num_frames:
@@ -175,8 +174,7 @@ class ParallelTiledVAE(ABC):
"""
Parallel version of tiled_decode that distributes both temporal and spatial computation across GPUs
"""
world_size, rank = get_sequence_model_parallel_world_size(
), get_sequence_model_parallel_rank()
world_size, rank = get_sp_world_size(), get_sp_parallel_rank()
B, C, T, H, W = z.shape
# Calculate parameters
+1
View File
@@ -1,3 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
# Copyright 2025 StepFun Inc. All Rights Reserved.
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
@@ -15,9 +15,8 @@ import torch
from fastvideo.v1.configs.pipelines import (PipelineConfig,
get_pipeline_config_cls_for_name)
from fastvideo.v1.distributed import (init_distributed_environment,
initialize_model_parallel,
model_parallel_is_initialized)
from fastvideo.v1.distributed import (
maybe_init_distributed_environment_and_model_parallel)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
@@ -49,7 +48,8 @@ class ComposedPipelineBase(ABC):
model_path: str,
fastvideo_args: FastVideoArgs,
config: Optional[Dict[str, Any]] = None,
required_config_modules: Optional[List[str]] = None):
required_config_modules: Optional[List[str]] = None,
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None):
"""
Initialize the pipeline. After __init__, the pipeline should be ready to
use. The pipeline should be stateless and not hold any batch state.
@@ -81,11 +81,12 @@ class ComposedPipelineBase(ABC):
else:
self.config = config
self.maybe_init_distributed_environment(fastvideo_args)
maybe_init_distributed_environment_and_model_parallel(
fastvideo_args.tp_size, fastvideo_args.sp_size)
# Load modules directly in initialization
logger.info("Loading pipeline modules...")
self.modules = self.load_modules(fastvideo_args)
self.modules = self.load_modules(fastvideo_args, loaded_modules)
if fastvideo_args.training_mode:
assert self.training_args is not None
@@ -118,7 +119,14 @@ class ComposedPipelineBase(ABC):
| PipelineConfig]] = None,
args: Optional[argparse.Namespace] = None,
required_config_modules: Optional[List[str]] = None,
loaded_modules: Optional[Dict[str,
torch.nn.Module]] = None,
**kwargs) -> "ComposedPipelineBase":
"""
Load a pipeline from a pretrained model.
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
"""
config = None
# 1. If users provide a pipeline config, it will override the default pipeline config
if isinstance(pipeline_config, PipelineConfig):
@@ -141,14 +149,9 @@ class ComposedPipelineBase(ABC):
config_args.update(kwargs)
if args is None or args.inference_mode:
fastvideo_args = FastVideoArgs(model_path=model_path,
device_str=device or "cuda" if
torch.cuda.is_available() else "cpu",
**config_args)
fastvideo_args = FastVideoArgs(model_path=model_path, **config_args)
fastvideo_args.model_path = model_path
fastvideo_args.device_str = device or "cuda" if torch.cuda.is_available(
) else "cpu"
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
else:
@@ -156,12 +159,9 @@ class ComposedPipelineBase(ABC):
fastvideo_args = TrainingArgs.from_cli_args(args)
# TODO(will): fix this so that its not so ugly
fastvideo_args.model_path = model_path
fastvideo_args.device_str = device or "cuda" if torch.cuda.is_available(
) else "cpu"
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
fastvideo_args.num_gpus = int(os.environ.get("WORLD_SIZE", 1))
fastvideo_args.use_cpu_offload = False
# make sure we are in training mode
fastvideo_args.inference_mode = False
@@ -173,38 +173,12 @@ class ComposedPipelineBase(ABC):
assert fastvideo_args.master_weight_type == 'fp32', 'only fp32 is supported for training'
# assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
fastvideo_args.check_fastvideo_args()
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
return cls(model_path,
fastvideo_args,
required_config_modules=required_config_modules)
def maybe_init_distributed_environment(self, fastvideo_args: FastVideoArgs):
if model_parallel_is_initialized():
return
local_rank = int(os.environ.get("LOCAL_RANK", -1))
world_size = int(os.environ.get("WORLD_SIZE", -1))
rank = int(os.environ.get("RANK", -1))
if local_rank == -1 or world_size == -1 or rank == -1:
raise ValueError(
"Local rank, world size, and rank must be set. Use torchrun to launch the script or pass rank to the worker process."
)
torch.cuda.set_device(local_rank)
init_distributed_environment(world_size=world_size,
rank=rank,
local_rank=local_rank)
assert fastvideo_args.tp_size is not None, "tp_size must be set"
assert fastvideo_args.sp_size is not None, "sp_size must be set"
initialize_model_parallel(
tensor_model_parallel_size=fastvideo_args.tp_size,
sequence_model_parallel_size=fastvideo_args.sp_size,
data_parallel_size=fastvideo_args.dp_size)
device = torch.device(f"cuda:{local_rank}")
fastvideo_args.device = device
required_config_modules=required_config_modules,
loaded_modules=loaded_modules)
def get_module(self, module_name: str, default_value: Any = None) -> Any:
if module_name not in self.modules:
@@ -266,9 +240,15 @@ class ComposedPipelineBase(ABC):
"""
return
def load_modules(self, fastvideo_args: FastVideoArgs) -> Dict[str, Any]:
def load_modules(
self,
fastvideo_args: FastVideoArgs,
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None
) -> Dict[str, Any]:
"""
Load the modules from the config.
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
"""
logger.info("Loading pipeline modules from config: %s", self.config)
modules_config = deepcopy(self.config)
@@ -297,6 +277,10 @@ class ComposedPipelineBase(ABC):
if module_name not in required_modules:
logger.info("Skipping module %s", module_name)
continue
if loaded_modules is not None and module_name in loaded_modules:
logger.info("Using module %s already provided", module_name)
modules[module_name] = loaded_modules[module_name]
continue
component_model_path = os.path.join(self.model_path, module_name)
module = PipelineComponentLoader.load_module(
module_name=module_name,

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