Compare commits

...
Author SHA1 Message Date
William Lin 97d4b984c9 [misc] [ci] fix e2e preprocess+training data path (#521) 2025-06-14 22:37:51 -07:00
Wenxuan Tan 2a8953d74d [Refactor] Fix attn backend selection not correctly setting env variable (#516) 2025-06-15 00:04:54 -05:00
Yongqi Chen 8801b10da7 [Bugfix][Preprocess]fix mini dataset name (#520) 2025-06-14 22:03:22 -07:00
William Lin 6b413f2ec4 [CI] [Training] drop negative prompt in validation dataset and CI test for preprocess + training overfit (#519) 2025-06-14 18:50:17 -07:00
Yongqi Chen 28b72694aa [Feature][Preprocess]Add Readme doc for preprocess (#518) 2025-06-14 20:41:13 -04:00
Yongqi Chen 4afb0cfe4f [Feature][Training]vsa for t2v training ready (#513) 2025-06-14 01:08:00 -04:00
Zhang Peiyuan 3eec1281cf [misc] Fix preprocessing and dataloader extra padding (#514) 2025-06-13 15:15:33 -07:00
Wenxuan Tan 0660489e38 [CI] Restrict training CI to v1 (#508) 2025-06-12 15:26:05 -07:00
Zhang Peiyuan dd871a17bf fix logging (#509) 2025-06-12 15:24:12 -07:00
dc11529862 [Refactor][Configurations] clean config orgnization (#505)
Co-authored-by: Peiyuan Zhang <a1286225768@gmail.com>
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2025-06-12 13:27:08 -07:00
Zhang Peiyuan ffabf85e31 [feat] Add parquet iterable dataset. (#506) 2025-06-12 04:30:56 -04:00
William Lin c0026ca5ba [CI] [Training] Initial e2e small training test (#504) 2025-06-11 13:53:36 -07:00
Zhang Peiyuan 0f2bbe71ac [misc] rename dp_size to hdsp_replicate_dim (#491) 2025-06-10 16:36:56 -07:00
Yongqi ChenandJerryZhou54 2a46902ecb [Feature][VSA]Update STA publish workflow (#498)
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
2025-06-10 19:33:34 -04:00
Zhang Peiyuan 66012d3a4c [Feat][Dataloader] 1/n Refactor parquet map-style dataloader (#492) 2025-06-10 16:00:13 -07:00
92 changed files with 2996 additions and 2002 deletions
+50
View File
@@ -39,6 +39,16 @@ on:
required: false
default: false
type: boolean
run_training_test:
description: "Run training-test"
required: false
default: false
type: boolean
run_nightly_test:
description: "Run nightly-test"
required: false
default: false
type: boolean
env:
PYTHONUNBUFFERED: "1"
@@ -59,6 +69,7 @@ jobs:
encoder-test: ${{ steps.filter.outputs.encoder-test }}
vae-test: ${{ steps.filter.outputs.vae-test }}
transformer-test: ${{ steps.filter.outputs.transformer-test }}
training-test: ${{ steps.filter.outputs.training-test }}
steps:
- uses: actions/checkout@v4
- uses: dorny/paths-filter@v3
@@ -79,6 +90,8 @@ jobs:
- 'fastvideo/v1/tests/transformers/**'
- 'fastvideo/v1/layers/**'
- 'fastvideo/v1/attention/**'
training-test:
- 'fastvideo/v1/**'
encoder-test:
needs: change-filter
@@ -160,6 +173,43 @@ jobs:
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
training-test:
needs: change-filter
if: >-
(github.event_name != 'workflow_dispatch' && github.event.pull_request.draft == false) ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_training_test == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "training-test"
gpu_type: "NVIDIA A40"
gpu_count: 4
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/training -srP"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
nightly-test:
if: >-
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "nightly-test"
gpu_type: "NVIDIA A40"
gpu_count: 4
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/${{ github.event.inputs.custom_image || 'fastvideo-dev:latest' }}"
test_command: "pip install -e .[test] && pytest ./fastvideo/v1/tests/nightly/test_e2e_overfit_single_sample.py -vs"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
runpod-cleanup:
needs: [encoder-test, vae-test, transformer-test, ssim-test] # Add other jobs to this list as you create them
+4 -1
View File
@@ -43,6 +43,8 @@ on:
required: true
RUNPOD_PRIVATE_KEY:
required: true
WANDB_API_KEY:
required: false
jobs:
run-test:
@@ -55,7 +57,7 @@ jobs:
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: "3.10"
python-version: "3.12"
- name: Set up SSH key
run: |
@@ -72,6 +74,7 @@ jobs:
JOB_ID: ${{ inputs.job_id }}
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
timeout-minutes: ${{ inputs.timeout_minutes }}
run: >-
python .github/scripts/runpod_api.py
+17 -13
View File
@@ -5,7 +5,7 @@ on:
branches:
- main
paths:
- "csrc/sliding_tile_attention/setup.py"
- "csrc/attn/setup_sta.py"
workflow_dispatch:
jobs:
@@ -23,13 +23,13 @@ jobs:
- name: Check if version changed
id: check-version
run: |
cd csrc/sliding_tile_attention
cd csrc/attn
# Get current commit's version
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup.py)
NEW_VERSION=$(grep -oP 'VERSION\s*=\s*"\K[^"]+' setup_sta.py)
echo "New version: $NEW_VERSION"
# Get previous version from git history
OLD_VERSION=$(git show HEAD~1:./setup.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
OLD_VERSION=$(git show HEAD~1:./setup_sta.py | grep -oP 'VERSION\s*=\s*"\K[^"]+' || echo "0.0.0")
echo "Old version: $OLD_VERSION"
if [ "$NEW_VERSION" != "$OLD_VERSION" ]; then
@@ -136,19 +136,21 @@ jobs:
- name: Build wheel
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
# However this still fails so I'm using a newer version of setuptools
pip install setuptools
pip install ninja packaging wheel
cd csrc/sliding_tile_attention # Move into the correct folder
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
python setup.py bdist_wheel --dist-dir=dist
cd csrc/attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup_sta.py bdist_wheel --dist-dir=dist
- name: Rename wheel file
run: |
cd csrc/sliding_tile_attention
cd csrc/attn
CUDA_SHORT_VERSION=$(echo ${{ matrix.cuda-version }} | cut -d. -f1,2 | sed 's/\.//g')
TORCH_SHORT_VERSION=$(echo ${{ matrix.torch-version }} | cut -d. -f1,2)
@@ -163,7 +165,7 @@ jobs:
uses: actions/upload-artifact@v4
with:
name: ${{ env.wheel_name }}
path: csrc/sliding_tile_attention/dist/*.whl
path: csrc/attn/dist/*.whl
retention-days: 90
publish_package:
@@ -229,17 +231,19 @@ jobs:
- name: Build source distribution
run: |
export PYTHONPATH=$GITHUB_WORKSPACE:$PYTHONPATH
# We want setuptools >= 49.6.0 otherwise we can't compile the extension if system CUDA version is 11.7 and pytorch cuda version is 11.6
# https://github.com/pytorch/pytorch/blob/664058fa83f1d8eede5d66418abff6e20bd76ca8/torch/utils/cpp_extension.py#L810
# However this still fails so I'm using a newer version of setuptools
pip install setuptools
pip install ninja packaging wheel
cd csrc/sliding_tile_attention # Move into the correct folder
git submodule update --init --recursive tk # Ensure ThunderKittens submodule is initialized
python setup.py sdist --dist-dir=dist
cd csrc/attn # Move into the correct folder
git submodule update --init --recursive # Ensure ThunderKittens submodule is initialized
python setup_sta.py sdist --dist-dir=dist
- name: Publish release distributions to PyPI
uses: pypa/gh-action-pypi-publish@release/v1
with:
packages-dir: csrc/sliding_tile_attention/dist/
packages-dir: csrc/attn/dist/
+1 -1
View File
@@ -28,4 +28,4 @@ jobs:
- name: Run Pytest
run: |
pytest --ignore csrc/sliding_tile_attention/test
pytest --ignore csrc/attn/test
-1
View File
@@ -27,7 +27,6 @@ env
**/build/
**.pyc
**.txt
csrc/attn/tk/
# Distribution / packaging
build/
+2 -2
View File
@@ -1,3 +1,3 @@
[submodule "csrc/sliding_tile_attention/tk"]
path = csrc/sliding_tile_attention/tk
[submodule "csrc/attn/tk"]
path = csrc/attn/tk
url = https://github.com/HazyResearch/ThunderKittens.git
+13 -4
View File
@@ -6,7 +6,6 @@
## Installation
We test our code on Pytorch 2.5.0 and CUDA>=12.4. Currently we only have implementation on H100.
First, install C++20 for ThunderKittens:
```bash
sudo apt update
sudo apt install gcc-11 g++-11
@@ -16,17 +15,27 @@ sudo update-alternatives --install /usr/bin/gcc gcc /usr/bin/gcc-11 100 --slave
sudo apt update
sudo apt install clang-11
```
Install STA:
## Environment Setup
First, set up your CUDA environment:
```bash
export CUDA_HOME=/usr/local/cuda-12.4
export PATH=${CUDA_HOME}/bin:${PATH}
export LD_LIBRARY_PATH=${CUDA_HOME}/lib64:$LD_LIBRARY_PATH
git submodule update --init --recursive
python setup.py install
```
## Install Sliding Tile Attention (STA)
```bash
python setup_sta.py install
```
## Install Video Sparse Attention (VSA)
```bash
python setup_vsa.py install
```
## Usage
```python
from st_attn import sliding_tile_attention
# assuming video size (T, H, W) = (30, 48, 80), text tokens = 256 with padding.
+2 -2
View File
@@ -2,8 +2,8 @@ import torch
import argparse
from flash_attn.utils.benchmark import benchmark_forward
from flash_attn import flash_attn_func
from st_attn import block_sparse_attention_fwd, block_sparse_attention_backward, BlockSparseAttentionFunction
from st_attn import BLOCK_M, BLOCK_N
from vsa import block_sparse_attention_fwd, block_sparse_attention_backward, BlockSparseAttentionFunction
from vsa import BLOCK_M, BLOCK_N
import numpy as np
import random
+4 -4
View File
@@ -9,14 +9,14 @@ flex_attention = torch.compile(flex_attention, dynamic=False)
def flex_test(Q, K, V, kernel_size):
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (36, 48, 48), 39, 'cuda', 0)
mask = get_sliding_tile_attention_mask(kernel_size, (6, 8, 8), (18, 48, 80), 0, 'cuda', 0)
output = flex_attention(Q, K, V, block_mask=mask)
return output
def h100_fwd_kernel_test(Q, K, V, kernel_size):
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 39, False)
o = sliding_tile_attention(Q, K, V, [kernel_size] * 24, 0, False, '18x48x80')
return o
@@ -37,7 +37,7 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
'max_diff': 0
},
}
kernel_size_ls = [(6, 1, 6), (6, 6, 1)]
kernel_size_ls = [(3, 3, 5), (3, 1, 10)]
from tqdm import tqdm
for kernel_size in tqdm(kernel_size_ls):
for _ in range(num_iterations):
@@ -74,7 +74,7 @@ def check_correctness(b, h, n, d, causal, mean, std, num_iterations=50, error_mo
# Example usage
b, h, d = 2, 24, 128
n = 82944 # Sequence length
n = 69120 # Sequence length
causal = False
mean = 1e-1
std = 10
Submodule
+1
Submodule csrc/attn/tk added at 1719fb7264
+284 -24
View File
@@ -2,7 +2,7 @@ import math
import torch
from torch.utils.checkpoint import detach_variable
from typing import Tuple
try:
from vsa_cuda import block_sparse_fwd, block_sparse_bwd
except ImportError:
@@ -12,33 +12,116 @@ except ImportError:
BLOCK_M = 64
BLOCK_N = 64
def block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num):
def video_sparse_attn(q, k, v, topk, block_size, compress_attn_weight=None):
"""
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks].
[*, *, i, j] = 1 means the i-th q block should attend to the j-th kv block.
q: [batch_size, num_heads, seq_len, head_dim]
k: [batch_size, num_heads, seq_len, head_dim]
v: [batch_size, num_heads, seq_len, head_dim]
topk: int
block_size: int or tuple of 3 ints
video_shape: tuple of (T, H, W)
compress_attn_weight: [batch_size, num_heads, seq_len, head_dim]
select_attn_weight: [batch_size, num_heads, seq_len, head_dim]
V1 of sparse attention. Include compress attn and sparse attn branch, use average pooling to compress.
Assume q, k, v is flattened in this way: [batch_size, num_heads, T//block_size[0], H//block_size[1], W//block_size[2], block_size[0], block_size[1], block_size[2]]
"""
# assert all elements in q2k_block_sparse_num can be devisible by 2
o, lse = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
return o, lse
def block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num):
grad_q, grad_k, grad_v = block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
return grad_q, grad_k, grad_v
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
class BlockSparseAttentionFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
o, lse = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
ctx.save_for_backward(q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num)
return o
block_elements = block_size[0] * block_size[1] * block_size[2]
assert block_elements % 64 == 0 and block_elements >= 64
assert q.shape[2] % block_elements == 0
batch_size, num_heads, seq_len, head_dim = q.shape
# compress attn
q_compress = q.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
k_compress = k.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
v_compress = v.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
@staticmethod
def backward(ctx, grad_output):
q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num = ctx.saved_tensors
grad_q, grad_k, grad_v = block_sparse_attention_backward(
q, k, v, o, lse, grad_output, k2q_block_sparse_index, k2q_block_sparse_num
)
return grad_q, grad_k, grad_v, None, None, None, None
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
v_compress)
output_compress = output_compress.view(batch_size, num_heads,
seq_len // block_elements, 1,
head_dim)
output_compress = output_compress.repeat(1, 1, 1, block_elements,
1).view(batch_size, num_heads,
seq_len, head_dim)
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num = generate_topk_block_sparse_pattern(
block_attn_score, topk)
output_select = block_sparse_attn(q, k, v, q2k_block_sparse_index,
q2k_block_sparse_num,
k2q_block_sparse_index,
k2q_block_sparse_num)
if compress_attn_weight is not None:
final_output = output_compress * compress_attn_weight + output_select
else:
final_output = output_compress + output_select
return final_output
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
QK = torch.matmul(q, k.transpose(-2, -1))
QK /= (q.size(-1)**0.5)
# Causal mask removed since causal is always false
QK = torch.nn.functional.softmax(QK, dim=-1)
output = torch.matmul(QK, v)
return output, QK
def generate_topk_block_sparse_pattern(block_attn_score: torch.Tensor,
topk: int):
"""
Generate a block sparse pattern where each q block attends to exactly topk kv blocks,
based on the provided attention scores.
Args:
block_attn_score: [bs, h, num_q_blocks, num_kv_blocks]
Attention scores between query and key blocks
topk: int
Number of kv blocks each q block attends to
Returns:
q2k_block_sparse_index: [bs, h, num_q_blocks, topk]
Contains the indices of kv blocks that each q block attends to.
q2k_block_sparse_num: [bs, h, num_q_blocks]
Contains the number of kv blocks that each q block attends to (all equal to topk).
k2q_block_sparse_index: [bs, h, num_kv_blocks, max_q_per_kv]
Contains the indices of q blocks that attend to each kv block.
k2q_block_sparse_num: [bs, h, num_kv_blocks]
Contains the number of q blocks that attend to each kv block.
"""
device = block_attn_score.device
# Extract dimensions from block_attn_score
bs, h, num_q_blocks, num_kv_blocks = block_attn_score.shape
sorted_result = torch.sort(block_attn_score, dim=-1, descending=True)
sorted_indice = sorted_result.indices
q2k_block_sparse_index, _ = torch.sort(sorted_indice[:, :, :, :topk],
dim=-1)
q2k_block_sparse_index = q2k_block_sparse_index.to(dtype=torch.int32)
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks),
topk,
device=device,
dtype=torch.int32)
block_map = topk_index_to_map(q2k_block_sparse_index,
num_kv_blocks,
transpose_map=True)
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(
block_map.transpose(2, 3))
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
@torch._dynamo.disable
def block_sparse_attn(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
@@ -61,6 +144,20 @@ def block_sparse_attn(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q
q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
)
def block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num):
"""
block_sparse_mask: [bs, h, num_q_blocks, num_kv_blocks].
[*, *, i, j] = 1 means the i-th q block should attend to the j-th kv block.
"""
# assert all elements in q2k_block_sparse_num can be devisible by 2
o, lse = block_sparse_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
return o, lse
def block_sparse_attention_backward(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num):
grad_output = grad_output.contiguous()
grad_q, grad_k, grad_v = block_sparse_bwd(q, k, v, o, l_vec, grad_output, k2q_block_sparse_index, k2q_block_sparse_num)
return grad_q, grad_k, grad_v
## pytorch sdpa version of block sparse ##
import triton
import triton.language as tl
@@ -128,6 +225,169 @@ def index_to_mask(q2k_block_sparse_index, q2k_block_sparse_num, BLOCK_Q, BLOCK_K
return mask
@triton.jit
def topk_index_to_map_kernel(
map_ptr,
index_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
topk: tl.constexpr,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
for i in tl.static_range(topk):
index = tl.load(index_ptr_base + i * index_kv_stride)
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
@triton.jit
def map_to_index_kernel(
map_ptr,
index_ptr,
index_num_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
index_num_bs_stride,
index_num_h_stride,
index_num_q_stride,
num_kv_blocks: tl.constexpr,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
num = 0
for i in tl.static_range(num_kv_blocks):
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
if map_entry:
tl.store(index_ptr_base + num * index_kv_stride, i)
num += 1
tl.store(
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
q * index_num_q_stride, num)
def topk_index_to_map(index: torch.Tensor,
num_kv_blocks: int,
transpose_map: bool = False):
"""
Convert topk indices to a map.
Args:
index: [bs, h, num_q_blocks, topk]
The topk indices tensor.
num_kv_blocks: int
The number of key-value blocks in the block_map returned
transpose_map: bool
If True, the block_map will be transposed on the final two dimensions.
Returns:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
A binary map where 1 indicates that the q block attends to the kv block.
"""
bs, h, num_q_blocks, topk = index.shape
if transpose_map is False:
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
dtype=torch.bool,
device=index.device)
else:
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
dtype=torch.bool,
device=index.device)
block_map = block_map.transpose(2, 3)
grid = (bs, h, num_q_blocks)
topk_index_to_map_kernel[grid](
block_map,
index,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
topk=topk,
)
return block_map
def map_to_index(block_map: torch.Tensor):
"""
Convert a block map to indices and counts.
Args:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
The block map tensor.
Returns:
index: [bs, h, num_q_blocks, num_kv_blocks]
The indices of the blocks.
index_num: [bs, h, num_q_blocks]
The number of blocks for each q block.
"""
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
index = torch.full((block_map.shape),
-1,
dtype=torch.int32,
device=block_map.device)
index_num = torch.empty((bs, h, num_q_blocks),
dtype=torch.int32,
device=block_map.device)
grid = (bs, h, num_q_blocks)
map_to_index_kernel[grid](
block_map,
index,
index_num,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
index_num.stride(0),
index_num.stride(1),
index_num.stride(2),
num_kv_blocks=num_kv_blocks,
)
return index, index_num
class BlockSparseAttentionFunction(torch.autograd.Function):
@staticmethod
def forward(ctx, q, k, v, q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num):
o, lse = block_sparse_attention_fwd(q, k, v, q2k_block_sparse_index, q2k_block_sparse_num)
ctx.save_for_backward(q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num)
return o
@staticmethod
def backward(ctx, grad_output):
q, k, v, o, lse, k2q_block_sparse_index, k2q_block_sparse_num = ctx.saved_tensors
grad_q, grad_k, grad_v = block_sparse_attention_backward(
q, k, v, o, lse, grad_output, k2q_block_sparse_index, k2q_block_sparse_num
)
return grad_q, grad_k, grad_v, None, None, None, None
class DummyOperator(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
@@ -207,4 +467,4 @@ class BlockSparseAttnTorch:
o = DummyOperator.apply(output)
o.register_hook(self.recompute)
return o
return o
+15 -45
View File
@@ -7,70 +7,40 @@ To save GPU memory, we precompute text embeddings and VAE latents to eliminate t
We provide a sample dataset to help you get started. Download the source media using the following command:
```bash
python scripts/huggingface/download_hf.py --repo_id=FastVideo/Image-Vid-Finetune-Src --local_dir=data/Image-Vid-Finetune-Src --repo_type=dataset
python scripts/huggingface/download_hf.py --repo_id=FastVideo/mini_i2v_dataset --local_dir=FastVideo/mini_i2v_dataset --repo_type=dataset
```
The folder `crush-smol_raw/` contains raw videos and captions for testing preprocessing, while `crush-smol_preprocessed/` contains latents prepared for testing training.
To preprocess the dataset for fine-tuning or distillation, run:
```
bash scripts/preprocess/preprocess_mochi_data.sh # for mochi
bash scripts/preprocess/preprocess_hunyuan_data.sh # for hunyuan
bash scripts/preprocess/v1_preprocess_wan_data_t2v # for wan
```
The preprocessed dataset will be stored in `Image-Vid-Finetune-Mochi` or `Image-Vid-Finetune-HunYuan` correspondingly.
## Process your own dataset
If you wish to create your own dataset for finetuning or distillation, please structure you video dataset in the following format:
If you wish to create your own dataset for finetuning or distillation, please refer `mini_i2v_dataset/crush-smol_raw/` to structure you video dataset in the following format:
```
path_to_dataset_folder/
├── media/
│ ├── 0.jpg
path_to_your_dataset_folder/
├── videos/
│ ├── 0.mp4
│ ├── 1.mp4
│ ├── 2.jpg
├── video2caption.json
└── merge.txt
├── videos.txt
└── prompt.txt
```
Format the JSON file as a list, where each item represents a media source:
To geranate the `videos2caption.json` and `merge.txt`, run
For image media,
```
{
"path": "0.jpg",
"cap": ["captions"]
}
``` python
python scripts/dataset_preparation/prepare_json_file.py --data_folder mini_i2v_dataset/crush-smol_raw/ --output your_output_folder
```
For video media,
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/v1_preprocess_****.sh` accordingly and run:
```
{
"path": "1.mp4",
"resolution": {
"width": 848,
"height": 480
},
"fps": 30.0,
"duration": 6.033333333333333,
"cap": [
"caption"
]
}
```
Use a txt file (merge.txt) to contain the source folder for media and the JSON file for meta information:
```
path_to_media_source_foder,path_to_json_file
```
Adjust the `DATA_MERGE_PATH` and `OUTPUT_DIR` in `scripts/preprocess/preprocess_****_data.sh` accordingly and run:
```
bash scripts/preprocess/preprocess_****_data.sh
bash scripts/preprocess/v1_preprocess_****.sh
```
The preprocessed data will be put into the `OUTPUT_DIR` and the `videos2caption.json` can be used in finetune and distill scripts.
@@ -104,13 +104,7 @@ if __name__ == "__main__":
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
args = parser.parse_args()
main(args)
-7
View File
@@ -671,13 +671,6 @@ if __name__ == "__main__":
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
-7
View File
@@ -693,13 +693,6 @@ if __name__ == "__main__":
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
+1 -7
View File
@@ -520,13 +520,7 @@ if __name__ == "__main__":
help=("Whether training should be resumed from a previous lora checkpoint. Use a path saved by"
' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.'),
)
parser.add_argument(
"--logging_dir",
type=str,
default="logs",
help=("[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to"
" *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***."),
)
# optimizer & scheduler & Training
parser.add_argument("--num_train_epochs", type=int, default=100)
@@ -287,16 +287,10 @@ class SlidingTileAttentionImpl(AttentionImpl):
forward_batch.mask_search_final_result_pos[timestep].append(
layer_loss_save)
else:
# windows = [
# self.mask_strategy[timestep][layer_idx][head_idx + start_head]
# for head_idx in range(head_num)
# ]
windows = [
STA_param[head_idx + start_head] for head_idx in range(head_num)
]
# if has_text is False:
# from IPython import embed
# embed()
hidden_states = sliding_tile_attention(
query, key, value, windows, text_length, has_text,
self.dit_seq_shape_str).transpose(1, 2)
@@ -1,18 +1,15 @@
# SPDX-License-Identifier: Apache-2.0
import math
from dataclasses import dataclass
from typing import Any, List, Optional, Type, cast
from typing import List, Optional, Type
import torch
import triton
import triton.language as tl
from einops import rearrange
try:
from vsa import block_sparse_attn
except ImportError: # noqa: E722
block_sparse_attn = None
from typing import Tuple
from vsa import video_sparse_attn
except ImportError:
video_sparse_attn = None
from fastvideo.v1.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
@@ -75,14 +72,18 @@ class VideoSparseAttentionMetadataBuilder(AttentionMetadataBuilder):
if forward_batch.latents is None:
raise ValueError("latents cannot be None")
raw_latent_shape = forward_batch.latents.shape
patch_size = fastvideo_args.dit_config.patch_size
raw_latent_shape = forward_batch.raw_latent_shape
if raw_latent_shape is None:
raise ValueError("raw_latent_shape cannot be None")
patch_size = fastvideo_args.pipeline_config.dit_config.patch_size
dit_seq_shape = [
raw_latent_shape[2] // patch_size[0],
raw_latent_shape[3] // patch_size[1],
raw_latent_shape[4] // patch_size[2]
]
VSA_sparsity = forward_batch.VSA_sparsity
return VideoSparseAttentionMetadata(current_timestep=current_timestep,
dit_seq_shape=dit_seq_shape,
VSA_sparsity=VSA_sparsity)
@@ -177,12 +178,16 @@ class VideoSparseAttentionImpl(AttentionImpl):
value = value.transpose(1, 2).contiguous()
gate_compress = gate_compress.transpose(1, 2).contiguous()
VSA_sparsity = attn_metadata.VSA_sparsity
cur_topk = math.ceil(
(1 - attn_metadata.VSA_sparsity) *
(1 - VSA_sparsity) *
(self.img_seq_length / math.prod(self.VSA_base_tile_size)))
# Cast to Any to bypass type checking for untyped function
hidden_states = cast(Any, sparse_attn_c_s_p)(
if video_sparse_attn is None:
raise NotImplementedError("video_sparse_attn is not installed")
hidden_states = video_sparse_attn(
query,
key,
value,
@@ -191,267 +196,3 @@ class VideoSparseAttentionImpl(AttentionImpl):
compress_attn_weight=gate_compress).transpose(1, 2)
return hidden_states
def torch_attention(q, k, v) -> Tuple[torch.Tensor, torch.Tensor]:
QK = torch.matmul(q, k.transpose(-2, -1))
QK /= (q.size(-1)**0.5)
# Causal mask removed since causal is always false
QK = torch.nn.functional.softmax(QK, dim=-1)
output = torch.matmul(QK, v)
return output, QK
def sparse_attn_c_s_p(q, k, v, topk, block_size, compress_attn_weight=None):
"""
q: [batch_size, num_heads, seq_len, head_dim]
k: [batch_size, num_heads, seq_len, head_dim]
v: [batch_size, num_heads, seq_len, head_dim]
topk: int
block_size: int or tuple of 3 ints
video_shape: tuple of (T, H, W)
compress_attn_weight: [batch_size, num_heads, seq_len, head_dim]
select_attn_weight: [batch_size, num_heads, seq_len, head_dim]
V1 of sparse attention. Include compress attn and sparse attn branch, use average pooling to compress.
Assume q, k, v is flattened in this way: [batch_size, num_heads, T//block_size[0], H//block_size[1], W//block_size[2], block_size[0], block_size[1], block_size[2]]
"""
if isinstance(block_size, int):
block_size = (block_size, block_size, block_size)
block_elements = block_size[0] * block_size[1] * block_size[2]
assert block_elements % 64 == 0 and block_elements >= 64
assert q.shape[2] % block_elements == 0
batch_size, num_heads, seq_len, head_dim = q.shape
# compress attn
q_compress = q.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
k_compress = k.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
v_compress = v.view(batch_size, num_heads, seq_len // block_elements,
block_elements, head_dim).mean(dim=3)
output_compress, block_attn_score = torch_attention(q_compress, k_compress,
v_compress)
output_compress = output_compress.view(batch_size, num_heads,
seq_len // block_elements, 1,
head_dim)
output_compress = output_compress.repeat(1, 1, 1, block_elements,
1).view(batch_size, num_heads,
seq_len, head_dim)
q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num = generate_topk_block_sparse_pattern(
block_attn_score, topk)
output_select = block_sparse_attn(q, k, v, q2k_block_sparse_index,
q2k_block_sparse_num,
k2q_block_sparse_index,
k2q_block_sparse_num)
if compress_attn_weight is not None:
final_output = output_compress * compress_attn_weight + output_select
else:
final_output = output_compress + output_select
return final_output
@triton.jit
def topk_index_to_map_kernel(
map_ptr,
index_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
topk: tl.constexpr,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
for i in tl.static_range(topk):
index = tl.load(index_ptr_base + i * index_kv_stride)
tl.store(map_ptr_base + index * map_kv_stride, 1.0)
@triton.jit
def map_to_index_kernel(
map_ptr,
index_ptr,
index_num_ptr,
map_bs_stride,
map_h_stride,
map_q_stride,
map_kv_stride,
index_bs_stride,
index_h_stride,
index_q_stride,
index_kv_stride,
index_num_bs_stride,
index_num_h_stride,
index_num_q_stride,
num_kv_blocks: tl.constexpr,
):
b, h, q = tl.program_id(0), tl.program_id(1), tl.program_id(2)
index_ptr_base = index_ptr + b * index_bs_stride + h * index_h_stride + q * index_q_stride
map_ptr_base = map_ptr + b * map_bs_stride + h * map_h_stride + q * map_q_stride
num = 0
for i in tl.static_range(num_kv_blocks):
map_entry = tl.load(map_ptr_base + i * map_kv_stride)
if map_entry:
tl.store(index_ptr_base + num * index_kv_stride, i)
num += 1
tl.store(
index_num_ptr + b * index_num_bs_stride + h * index_num_h_stride +
q * index_num_q_stride, num)
def topk_index_to_map(index: torch.Tensor,
num_kv_blocks: int,
transpose_map: bool = False):
"""
Convert topk indices to a map.
Args:
index: [bs, h, num_q_blocks, topk]
The topk indices tensor.
num_kv_blocks: int
The number of key-value blocks in the block_map returned
transpose_map: bool
If True, the block_map will be transposed on the final two dimensions.
Returns:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
A binary map where 1 indicates that the q block attends to the kv block.
"""
bs, h, num_q_blocks, topk = index.shape
if transpose_map is False:
block_map = torch.zeros((bs, h, num_q_blocks, num_kv_blocks),
dtype=torch.bool,
device=index.device)
else:
block_map = torch.zeros((bs, h, num_kv_blocks, num_q_blocks),
dtype=torch.bool,
device=index.device)
block_map = block_map.transpose(2, 3)
grid = (bs, h, num_q_blocks)
topk_index_to_map_kernel[grid](
block_map,
index,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
topk=topk,
)
return block_map
def map_to_index(block_map: torch.Tensor):
"""
Convert a block map to indices and counts.
Args:
block_map: [bs, h, num_q_blocks, num_kv_blocks]
The block map tensor.
Returns:
index: [bs, h, num_q_blocks, num_kv_blocks]
The indices of the blocks.
index_num: [bs, h, num_q_blocks]
The number of blocks for each q block.
"""
bs, h, num_q_blocks, num_kv_blocks = block_map.shape
index = torch.full((block_map.shape),
-1,
dtype=torch.int32,
device=block_map.device)
index_num = torch.empty((bs, h, num_q_blocks),
dtype=torch.int32,
device=block_map.device)
grid = (bs, h, num_q_blocks)
map_to_index_kernel[grid](
block_map,
index,
index_num,
block_map.stride(0),
block_map.stride(1),
block_map.stride(2),
block_map.stride(3),
index.stride(0),
index.stride(1),
index.stride(2),
index.stride(3),
index_num.stride(0),
index_num.stride(1),
index_num.stride(2),
num_kv_blocks=num_kv_blocks,
)
return index, index_num
def generate_topk_block_sparse_pattern(block_attn_score: torch.Tensor,
topk: int):
"""
Generate a block sparse pattern where each q block attends to exactly topk kv blocks,
based on the provided attention scores.
Args:
block_attn_score: [bs, h, num_q_blocks, num_kv_blocks]
Attention scores between query and key blocks
topk: int
Number of kv blocks each q block attends to
Returns:
q2k_block_sparse_index: [bs, h, num_q_blocks, topk]
Contains the indices of kv blocks that each q block attends to.
q2k_block_sparse_num: [bs, h, num_q_blocks]
Contains the number of kv blocks that each q block attends to (all equal to topk).
k2q_block_sparse_index: [bs, h, num_kv_blocks, max_q_per_kv]
Contains the indices of q blocks that attend to each kv block.
k2q_block_sparse_num: [bs, h, num_kv_blocks]
Contains the number of q blocks that attend to each kv block.
"""
device = block_attn_score.device
# Extract dimensions from block_attn_score
bs, h, num_q_blocks, num_kv_blocks = block_attn_score.shape
sorted_result = torch.sort(block_attn_score, dim=-1, descending=True)
sorted_indice = sorted_result.indices
q2k_block_sparse_index, _ = torch.sort(sorted_indice[:, :, :, :topk],
dim=-1)
q2k_block_sparse_index = q2k_block_sparse_index.to(dtype=torch.int32)
q2k_block_sparse_num = torch.full((bs, h, num_q_blocks),
topk,
device=device,
dtype=torch.int32)
block_map = topk_index_to_map(q2k_block_sparse_index,
num_kv_blocks,
transpose_map=True)
k2q_block_sparse_index, k2q_block_sparse_num = map_to_index(
block_map.transpose(2, 3))
return q2k_block_sparse_index, q2k_block_sparse_num, k2q_block_sparse_index, k2q_block_sparse_num
+26 -27
View File
@@ -12,7 +12,7 @@ from fastvideo.v1.distributed.communication_op import (
from fastvideo.v1.distributed.parallel_state import (get_sp_parallel_rank,
get_sp_world_size)
from fastvideo.v1.forward_context import ForwardContext, get_forward_context
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
from fastvideo.v1.utils import get_compute_dtype
@@ -26,8 +26,8 @@ class DistributedAttention(nn.Module):
num_kv_heads: Optional[int] = None,
softmax_scale: Optional[float] = None,
causal: bool = False,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
supported_attention_backends: Optional[Tuple[
AttentionBackendEnum, ...]] = None,
prefix: str = "",
**extra_impl_args) -> None:
super().__init__()
@@ -45,13 +45,13 @@ class DistributedAttention(nn.Module):
dtype,
supported_attention_backends=supported_attention_backends)
impl_cls = attn_backend.get_impl_cls()
self.impl = impl_cls(num_heads=num_heads,
head_size=head_size,
causal=causal,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
prefix=f"{prefix}.impl",
**extra_impl_args)
self.attn_impl = impl_cls(num_heads=num_heads,
head_size=head_size,
causal=causal,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
prefix=f"{prefix}.impl",
**extra_impl_args)
self.num_heads = num_heads
self.head_size = head_size
self.num_kv_heads = num_kv_heads
@@ -100,7 +100,7 @@ class DistributedAttention(nn.Module):
scatter_dim=2,
gather_dim=1)
# Apply backend-specific preprocess_qkv
qkv = self.impl.preprocess_qkv(qkv, ctx_attn_metadata)
qkv = self.attn_impl.preprocess_qkv(qkv, ctx_attn_metadata)
# Concatenate with replicated QKV if provided
if replicated_q is not None:
@@ -116,7 +116,7 @@ class DistributedAttention(nn.Module):
q, k, v = qkv.chunk(3, dim=0)
output = self.impl.forward(q, k, v, ctx_attn_metadata)
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
# Redistribute back if using sequence parallelism
replicated_output = None
@@ -127,7 +127,7 @@ class DistributedAttention(nn.Module):
replicated_output = sequence_model_parallel_all_gather(
replicated_output.contiguous(), dim=2)
# Apply backend-specific postprocess_output
output = self.impl.postprocess_output(output, ctx_attn_metadata)
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
output = sequence_model_parallel_all_to_all_4D(output,
scatter_dim=1,
@@ -183,18 +183,17 @@ class DistributedAttention_VSA(DistributedAttention):
scatter_dim=2,
gather_dim=1)
qkvg = self.impl.preprocess_qkv(
qkvg, ctx_attn_metadata) # (yongqi) pass latent shape here?
qkvg = self.attn_impl.preprocess_qkv(qkvg, ctx_attn_metadata)
q, k, v, gate_compress = qkvg.chunk(4, dim=0)
output = self.impl.forward(q, k, v, gate_compress,
ctx_attn_metadata) # type: ignore[call-arg]
output = self.attn_impl.forward(
q, k, v, gate_compress, ctx_attn_metadata) # type: ignore[call-arg]
# Redistribute back if using sequence parallelism
replicated_output = None
# Apply backend-specific postprocess_output
output = self.impl.postprocess_output(output, ctx_attn_metadata)
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
output = sequence_model_parallel_all_to_all_4D(output,
scatter_dim=1,
@@ -212,8 +211,8 @@ class LocalAttention(nn.Module):
num_kv_heads: Optional[int] = None,
softmax_scale: Optional[float] = None,
causal: bool = False,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
supported_attention_backends: Optional[Tuple[
AttentionBackendEnum, ...]] = None,
**extra_impl_args) -> None:
super().__init__()
if softmax_scale is None:
@@ -229,12 +228,12 @@ class LocalAttention(nn.Module):
dtype,
supported_attention_backends=supported_attention_backends)
impl_cls = attn_backend.get_impl_cls()
self.impl = impl_cls(num_heads=num_heads,
head_size=head_size,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
causal=causal,
**extra_impl_args)
self.attn_impl = impl_cls(num_heads=num_heads,
head_size=head_size,
softmax_scale=self.softmax_scale,
num_kv_heads=num_kv_heads,
causal=causal,
**extra_impl_args)
self.num_heads = num_heads
self.head_size = head_size
self.num_kv_heads = num_kv_heads
@@ -265,5 +264,5 @@ class LocalAttention(nn.Module):
forward_context: ForwardContext = get_forward_context()
ctx_attn_metadata = forward_context.attn_metadata
output = self.impl.forward(q, k, v, ctx_attn_metadata)
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
return output
+14 -11
View File
@@ -11,13 +11,13 @@ import torch
import fastvideo.v1.envs as envs
from fastvideo.v1.attention.backends.abstract import AttentionBackend
from fastvideo.v1.logger import init_logger
from fastvideo.v1.platforms import _Backend, current_platform
from fastvideo.v1.platforms import AttentionBackendEnum, current_platform
from fastvideo.v1.utils import STR_BACKEND_ENV_VAR, resolve_obj_by_qualname
logger = init_logger(__name__)
def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
def backend_name_to_enum(backend_name: str) -> Optional[AttentionBackendEnum]:
"""
Convert a string backend name to a _Backend enum value.
@@ -27,11 +27,11 @@ def backend_name_to_enum(backend_name: str) -> Optional[_Backend]:
loaded.
"""
assert backend_name is not None
return _Backend[backend_name] if backend_name in _Backend.__members__ else \
return AttentionBackendEnum[backend_name] if backend_name in AttentionBackendEnum.__members__ else \
None
def get_env_variable_attn_backend() -> Optional[_Backend]:
def get_env_variable_attn_backend() -> Optional[AttentionBackendEnum]:
'''
Get the backend override specified by the FastVideo attention
backend environment variable, if one is specified.
@@ -53,10 +53,11 @@ def get_env_variable_attn_backend() -> Optional[_Backend]:
#
# THIS SELECTION TAKES PRECEDENCE OVER THE
# FASTVIDEO ATTENTION BACKEND ENVIRONMENT VARIABLE
forced_attn_backend: Optional[_Backend] = None
forced_attn_backend: Optional[AttentionBackendEnum] = None
def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
def global_force_attn_backend(
attn_backend: Optional[AttentionBackendEnum]) -> None:
'''
Force all attention operations to use a specified backend.
@@ -71,7 +72,7 @@ def global_force_attn_backend(attn_backend: Optional[_Backend]) -> None:
forced_attn_backend = attn_backend
def get_global_forced_attn_backend() -> Optional[_Backend]:
def get_global_forced_attn_backend() -> Optional[AttentionBackendEnum]:
'''
Get the currently-forced choice of attention backend,
or None if auto-selection is currently enabled.
@@ -82,7 +83,8 @@ def get_global_forced_attn_backend() -> Optional[_Backend]:
def get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
) -> Type[AttentionBackend]:
return _cached_get_attn_backend(head_size, dtype,
supported_attention_backends)
@@ -92,7 +94,8 @@ def get_attn_backend(
def _cached_get_attn_backend(
head_size: int,
dtype: torch.dtype,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
) -> Type[AttentionBackend]:
# Check whether a particular choice of backend was
# previously forced.
@@ -102,7 +105,7 @@ def _cached_get_attn_backend(
if not supported_attention_backends:
raise ValueError("supported_attention_backends is empty")
selected_backend = None
backend_by_global_setting: Optional[_Backend] = (
backend_by_global_setting: Optional[AttentionBackendEnum] = (
get_global_forced_attn_backend())
if backend_by_global_setting is not None:
selected_backend = backend_by_global_setting
@@ -125,7 +128,7 @@ def _cached_get_attn_backend(
@contextmanager
def global_force_attn_backend_context_manager(
attn_backend: _Backend) -> Generator[None, None, None]:
attn_backend: AttentionBackendEnum) -> Generator[None, None, None]:
'''
Globally force a FastVideo attention backend override within a
context manager, reverting the global attention backend
+5 -7
View File
@@ -4,7 +4,7 @@ from typing import Any, List, Optional, Tuple
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.layers.quantization import QuantizationConfig
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
@dataclass
@@ -13,12 +13,10 @@ class DiTArchConfig(ArchConfig):
_compile_conditions: list = field(default_factory=list)
_param_names_mapping: dict = field(default_factory=dict)
_lora_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: Tuple[_Backend,
...] = (_Backend.SLIDING_TILE_ATTN,
_Backend.SAGE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA,
_Backend.VIDEO_SPARSE_ATTN)
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.SLIDING_TILE_ATTN, AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN)
hidden_size: int = 0
num_attention_heads: int = 0
+3 -3
View File
@@ -6,14 +6,14 @@ import torch
from fastvideo.v1.configs.models.base import ArchConfig, ModelConfig
from fastvideo.v1.layers.quantization import QuantizationConfig
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
@dataclass
class EncoderArchConfig(ArchConfig):
architectures: List[str] = field(default_factory=lambda: [])
_supported_attention_backends: Tuple[_Backend, ...] = (_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)
_supported_attention_backends: Tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA)
output_hidden_states: bool = False
use_return_dict: bool = True
+11
View File
@@ -1,4 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import dataclasses
from dataclasses import dataclass, field
from typing import Any, Union
@@ -129,3 +131,12 @@ class VAEConfig(ModelConfig):
)
return parser
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "VAEConfig":
kwargs = {}
for attr in dataclasses.fields(cls):
value = getattr(args, attr.name, None)
if value is not None:
kwargs[attr.name] = value
return cls(**kwargs)
+2 -2
View File
@@ -3,7 +3,7 @@ from fastvideo.v1.configs.pipelines.base import (PipelineConfig,
from fastvideo.v1.configs.pipelines.hunyuan import (FastHunyuanConfig,
HunyuanConfig)
from fastvideo.v1.configs.pipelines.registry import (
get_pipeline_config_cls_for_name)
get_pipeline_config_cls_from_name)
from fastvideo.v1.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.v1.configs.pipelines.wan import (WanI2V480PConfig,
WanI2V720PConfig,
@@ -14,5 +14,5 @@ __all__ = [
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
"get_pipeline_config_cls_for_name"
"get_pipeline_config_cls_from_name"
]
+243 -18
View File
@@ -1,15 +1,17 @@
# SPDX-License-Identifier: Apache-2.0
import json
from dataclasses import asdict, dataclass, field, fields
from typing import Any, Callable, Dict, Optional, Tuple, cast
from typing import Any, Callable, Dict, List, Optional, Tuple, Union, cast
import torch
from fastvideo.v1.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
VAEConfig)
from fastvideo.v1.configs.models.encoders import BaseEncoderOutput
from fastvideo.v1.configs.utils import update_config_from_args
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import shallow_asdict
from fastvideo.v1.utils import (FlexibleArgumentParser, StoreBoolean,
shallow_asdict)
logger = init_logger(__name__)
@@ -22,59 +24,282 @@ def postprocess_text(output: BaseEncoderOutput) -> torch.tensor:
raise NotImplementedError
# config for a single pipeline
@dataclass
class PipelineConfig:
"""Base configuration for all pipeline architectures."""
model_path: str = ""
pipeline_config_path: Optional[str] = None
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
disable_autocast: bool = False
# Model configuration
precision: str = "bf16"
dit_config: DiTConfig = field(default_factory=DiTConfig)
dit_precision: str = "bf16"
# VAE configuration
vae_config: VAEConfig = field(default_factory=VAEConfig)
vae_precision: str = "fp16"
vae_tiling: bool = True
vae_sp: bool = True
vae_config: VAEConfig = field(default_factory=VAEConfig)
# DiT configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
# Image encoder configuration
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
image_encoder_precision: str = "fp32"
# Text encoder configuration
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp16", ))
DEFAULT_TEXT_ENCODER_PRECISIONS = ("fp16", )
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (EncoderConfig(), ))
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp16", ))
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: Tuple[Callable[[BaseEncoderOutput], torch.tensor],
...] = field(default_factory=lambda:
(postprocess_text, ))
# STA (Spatial-Temporal Attention) parameters
# LoRA parameters
lora_path: Optional[str] = None
lora_nickname: Optional[
str] = "default" # for swapping adapters in the pipeline
lora_target_names: Optional[List[
str]] = None # can restrict list of layers to adapt, e.g. ["q_proj"]
# StepVideo specific parameters
pos_magic: Optional[str] = None
neg_magic: Optional[str] = None
timesteps_scale: Optional[bool] = None
# STA (Sliding Tile Attention) parameters
mask_strategy_file_path: Optional[str] = None
STA_mode: str = "STA_inference"
STA_mode: Optional[str] = None
skip_time_steps: int = 15
# Compilation
enable_torch_compile: bool = False
# enable_torch_compile: bool = False
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser,
prefix: str = "") -> FlexibleArgumentParser:
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
# model_path will be conflicting with the model_path in FastVideoArgs,
# so we add it separately if prefix is not empty
if prefix_with_dot != "":
parser.add_argument(
f"--{prefix_with_dot}model-path",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}model_path",
default=PipelineConfig.model_path,
help="Path to the pretrained model",
)
parser.add_argument(
f"--{prefix_with_dot}pipeline-config-path",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}pipeline_config_path",
default=PipelineConfig.pipeline_config_path,
help="Path to the pipeline config",
)
parser.add_argument(
f"--{prefix_with_dot}embedded-cfg-scale",
type=float,
dest=f"{prefix_with_dot.replace('-', '_')}embedded_cfg_scale",
default=PipelineConfig.embedded_cfg_scale,
help="Embedded CFG scale",
)
parser.add_argument(
f"--{prefix_with_dot}flow-shift",
type=float,
dest=f"{prefix_with_dot.replace('-', '_')}flow_shift",
default=PipelineConfig.flow_shift,
help="Flow shift parameter",
)
# DiT configuration
parser.add_argument(
f"--{prefix_with_dot}dit-precision",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}dit_precision",
default=PipelineConfig.dit_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for the DiT model",
)
# VAE configuration
parser.add_argument(
f"--{prefix_with_dot}vae-precision",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}vae_precision",
default=PipelineConfig.vae_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for VAE",
)
parser.add_argument(
f"--{prefix_with_dot}vae-tiling",
action=StoreBoolean,
dest=f"{prefix_with_dot.replace('-', '_')}vae_tiling",
default=PipelineConfig.vae_tiling,
help="Enable VAE tiling",
)
parser.add_argument(
f"--{prefix_with_dot}vae-sp",
action=StoreBoolean,
dest=f"{prefix_with_dot.replace('-', '_')}vae_sp",
help="Enable VAE spatial parallelism",
)
# Text encoder configuration
parser.add_argument(
f"--{prefix_with_dot}text-encoder-precisions",
nargs="+",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}text_encoder_precisions",
default=PipelineConfig.DEFAULT_TEXT_ENCODER_PRECISIONS,
choices=["fp32", "fp16", "bf16"],
help="Precision for each text encoder",
)
# Image encoder configuration
parser.add_argument(
f"--{prefix_with_dot}image-encoder-precision",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}image_encoder_precision",
default=PipelineConfig.image_encoder_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for image encoder",
)
parser.add_argument(
f"--{prefix_with_dot}pos_magic",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}pos_magic",
default=PipelineConfig.pos_magic,
help="Positive magic prompt for sampling, used in stepvideo",
)
parser.add_argument(
f"--{prefix_with_dot}neg_magic",
type=str,
dest=f"{prefix_with_dot.replace('-', '_')}neg_magic",
default=PipelineConfig.neg_magic,
help="Negative magic prompt for sampling, used in stepvideo",
)
parser.add_argument(
f"--{prefix_with_dot}timesteps_scale",
type=bool,
dest=f"{prefix_with_dot.replace('-', '_')}timesteps_scale",
default=PipelineConfig.timesteps_scale,
help=
"Bool for applying scheduler scale in set_timesteps, used in stepvideo",
)
# Add VAE configuration arguments
from fastvideo.v1.configs.models.vaes.base import VAEConfig
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
# Add DiT configuration arguments
from fastvideo.v1.configs.models.dits.base import DiTConfig
DiTConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}dit-config")
return parser
def update_config_from_dict(self,
args: Dict[str, Any],
prefix: str = "") -> None:
prefix_with_dot = f"{prefix}." if (prefix.strip() != "") else ""
update_config_from_args(self, args, prefix, pop_args=True)
update_config_from_args(self.vae_config,
args,
f"{prefix_with_dot}vae_config",
pop_args=True)
update_config_from_args(self.dit_config,
args,
f"{prefix_with_dot}dit_config",
pop_args=True)
@classmethod
def from_pretrained(cls, model_path: str) -> "PipelineConfig":
"""
use the pipeline class setting from model_path to match the pipeline config
"""
from fastvideo.v1.configs.pipelines.registry import (
get_pipeline_config_cls_for_name)
pipeline_config_cls = get_pipeline_config_cls_for_name(model_path)
if pipeline_config_cls is not None:
pipeline_config = pipeline_config_cls()
else:
get_pipeline_config_cls_from_name)
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
return cast(PipelineConfig, pipeline_config_cls(model_path=model_path))
@classmethod
def from_kwargs(cls,
kwargs: Dict[str, Any],
config_cli_prefix: str = "") -> "PipelineConfig":
"""
Load PipelineConfig from kwargs Dictionary.
kwargs: dictionary of kwargs
config_cli_prefix: prefix of CLI arguments for this PipelineConfig instance
"""
from fastvideo.v1.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
prefix_with_dot = f"{config_cli_prefix}." if (config_cli_prefix.strip()
!= "") else ""
model_path: Optional[str] = kwargs.get(prefix_with_dot + 'model_path',
None) or kwargs.get('model_path')
pipeline_config_or_path: Optional[Union[str, PipelineConfig, Dict[
str, Any]]] = kwargs.get(prefix_with_dot + 'pipeline_config',
None) or kwargs.get('pipeline_config')
if model_path is None:
raise ValueError("model_path is required in kwargs")
# 1. Get the pipeline config class from the registry
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
# 2. Instantiate PipelineConfig
if pipeline_config_cls is None:
logger.warning(
"Couldn't find an optimal sampling param for %s. Using the default sampling param.",
"Couldn't find pipeline config for %s. Using the default pipeline config.",
model_path)
pipeline_config = cls()
else:
pipeline_config = pipeline_config_cls()
return cast(PipelineConfig, pipeline_config)
# 3. Load PipelineConfig from a json file or a PipelineConfig object if provided
if isinstance(pipeline_config_or_path, str):
pipeline_config.load_from_json(pipeline_config_or_path)
kwargs[prefix_with_dot +
'pipeline_config_path'] = pipeline_config_or_path
elif isinstance(pipeline_config_or_path, PipelineConfig):
pipeline_config = pipeline_config_or_path
elif isinstance(pipeline_config_or_path, dict):
pipeline_config.update_pipeline_config(pipeline_config_or_path)
# 4. Update PipelineConfig from CLI arguments if provided
kwargs[prefix_with_dot + 'model_path'] = model_path
pipeline_config.update_config_from_dict(kwargs, config_cli_prefix)
return pipeline_config
def check_pipeline_config(self) -> None:
if self.vae_sp and not self.vae_tiling:
raise ValueError(
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
)
if len(self.text_encoder_configs) != len(self.text_encoder_precisions):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
)
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
if len(self.preprocess_text_funcs) != len(self.postprocess_text_funcs):
raise ValueError(
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
def dump_to_json(self, file_path: str):
output_dict = shallow_asdict(self)
+1 -1
View File
@@ -80,7 +80,7 @@ class HunyuanConfig(PipelineConfig):
(llama_postprocess_text, clip_postprocess_text))
# Precision for each component
precision: str = "bf16"
dit_precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: ("fp16", "fp16"))
+63 -26
View File
@@ -19,7 +19,7 @@ from fastvideo.v1.utils import (maybe_download_model_index,
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
WEIGHT_CONFIG_REGISTRY: Dict[str, Type[PipelineConfig]] = {
PIPE_NAME_TO_CONFIG: Dict[str, Type[PipelineConfig]] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
@@ -51,37 +51,74 @@ PIPELINE_FALLBACK_CONFIG: Dict[str, Type[PipelineConfig]] = {
}
def get_pipeline_config_cls_for_name(
pipeline_name_or_path: str) -> Optional[type[PipelineConfig]]:
"""Get the appropriate config class for specific pretrained weights."""
def get_pipeline_config_cls_from_name(
pipeline_name_or_path: str) -> Type[PipelineConfig]:
"""Get the appropriate configuration class for a given pipeline name or path.
if os.path.exists(pipeline_name_or_path):
config = verify_model_config_and_directory(pipeline_name_or_path)
logger.warning(
"FastVideo may not correctly identify the optimal config for this model, as the local directory may have been renamed."
)
else:
config = maybe_download_model_index(pipeline_name_or_path)
This function implements a multi-step lookup process to find the most suitable
configuration class for a given pipeline. It follows this order:
1. Exact match in the PIPE_NAME_TO_CONFIG
2. Partial match in the PIPE_NAME_TO_CONFIG
3. Fallback to class name in the model_index.json
4. else raise an error
pipeline_name = config["_class_name"]
Args:
pipeline_name_or_path (str): The name or path of the pipeline. This can be:
- A registered model ID (e.g., "FastVideo/FastHunyuan-diffusers")
- A local path to a model directory
- A model ID that will be downloaded
Returns:
Type[PipelineConfig]: The configuration class that best matches the pipeline.
This will be one of:
- A specific weight configuration class if an exact match is found
- A fallback configuration class based on the pipeline architecture
- The base PipelineConfig class if no matches are found
Note:
- For local paths, the function will verify the model configuration
- For remote models, it will attempt to download the model index
- Warning messages are logged when falling back to less specific configurations
"""
pipeline_config_cls: Optional[Type[PipelineConfig]] = None
# First try exact match for specific weights
if pipeline_name_or_path in WEIGHT_CONFIG_REGISTRY:
return WEIGHT_CONFIG_REGISTRY[pipeline_name_or_path]
if pipeline_name_or_path in PIPE_NAME_TO_CONFIG:
pipeline_config_cls = PIPE_NAME_TO_CONFIG[pipeline_name_or_path]
# Try partial matches (for local paths that might include the weight ID)
for registered_id, config_class in WEIGHT_CONFIG_REGISTRY.items():
for registered_id, config_class in PIPE_NAME_TO_CONFIG.items():
if registered_id in pipeline_name_or_path:
return config_class
# If no match, try to use the fallback config
fallback_config = None
# Try to determine pipeline architecture for fallback
for pipeline_type, detector in PIPELINE_DETECTOR.items():
if detector(pipeline_name.lower()):
fallback_config = PIPELINE_FALLBACK_CONFIG.get(pipeline_type)
pipeline_config_cls = config_class
break
logger.warning("No match found for pipeline %s, using fallback config %s.",
pipeline_name_or_path, fallback_config)
return fallback_config
# If no match, try to use the fallback config
if pipeline_config_cls is None:
if os.path.exists(pipeline_name_or_path):
config = verify_model_config_and_directory(pipeline_name_or_path)
else:
config = maybe_download_model_index(pipeline_name_or_path)
logger.warning(
"Trying to use the config from the model_index.json. FastVideo may not correctly identify the optimal config for this model in this situation."
)
pipeline_name = config["_class_name"]
# Try to determine pipeline architecture for fallback
for pipeline_type, detector in PIPELINE_DETECTOR.items():
if detector(pipeline_name.lower()):
pipeline_config_cls = PIPELINE_FALLBACK_CONFIG.get(
pipeline_type)
break
if pipeline_config_cls is not None:
logger.warning(
"No match found for pipeline %s, using fallback config %s.",
pipeline_name_or_path, pipeline_config_cls)
if pipeline_config_cls is None:
raise ValueError(
f"No match found for pipeline {pipeline_name_or_path}, please check the pipeline name or path."
)
return pipeline_config_cls
-7
View File
@@ -39,7 +39,6 @@ class SamplingParam:
num_inference_steps: int = 50
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
VSA_sparsity: float = 0.0
# TeaCache parameters
enable_teacache: bool = False
@@ -185,12 +184,6 @@ class SamplingParam:
default=SamplingParam.image_path,
help="Path to input image for image-to-video generation",
)
parser.add_argument(
"--VSA-sparsity",
type=float,
default=SamplingParam.VSA_sparsity,
help="VSA attention sparsity",
)
return parser
+45
View File
@@ -0,0 +1,45 @@
from typing import Any, Dict
def update_config_from_args(config: Any,
args_dict: Dict[str, Any],
prefix: str = "",
pop_args: bool = False) -> None:
"""
Update configuration object from arguments dictionary.
Args:
config: The configuration object to update
args_dict: Dictionary containing arguments
prefix: Prefix for the configuration parameters in the args_dict.
If None, assumes direct attribute mapping without prefix.
"""
# Handle top-level attributes (no prefix)
args_not_to_remove = [
'model_path',
]
args_to_remove = []
if prefix.strip() == "":
for key, value in args_dict.items():
if hasattr(config, key) and value is not None:
if key == "text_encoder_precisions" and isinstance(value, list):
setattr(config, key, tuple(value))
else:
setattr(config, key, value)
if pop_args:
args_to_remove.append(key)
else:
# Handle nested attributes with prefix
prefix_with_dot = f"{prefix}."
for key, value in args_dict.items():
if key.startswith(prefix_with_dot) and value is not None:
attr_name = key[len(prefix_with_dot):]
if hasattr(config, attr_name):
setattr(config, attr_name, value)
if pop_args:
args_to_remove.append(key)
if pop_args:
for key in args_to_remove:
if key not in args_not_to_remove:
args_dict.pop(key)
+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,185 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import os
import pathlib
import time
import torch.distributed as dist
import torch.distributed.checkpoint as dist_cp
from fastvideo.v1.dataset.parquet_dataset_iterable_style import (
build_parquet_iterable_style_dataloader)
from fastvideo.v1.distributed import get_world_rank
from fastvideo.v1.distributed.parallel_state import (
cleanup_dist_env_and_memory, get_torch_device,
maybe_init_distributed_environment_and_model_parallel)
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
def main() -> None:
parser = argparse.ArgumentParser(
description="Benchmark parquet iterable style dataset loading speed")
parser.add_argument(
"--path",
type=str,
help="Path to parquet dataset",
)
parser.add_argument("--batch_size",
type=int,
default=4,
help="Batch size for DataLoader")
parser.add_argument("--num_data_workers",
type=int,
help="Number of DataLoader workers")
parser.add_argument("--num_epoch",
type=int,
default=2,
help="Number of epoches to benchmark")
parser.add_argument("--verify_resume",
action="store_true",
help="Verify resume")
parser.add_argument(
"--num_batches_per_epoch",
type=int,
default=1000,
help="Number of batches to benchmark",
)
parser.add_argument('--checkpoint_path',
type=str,
default='dataloader_checkpoint',
help='Path to save/load checkpoint')
'''
example launch command:
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 2 --num_epoch 2 --num_batches_per_epoch 2 --verify_resume
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5 --verify_resume
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_iterable_style.py --path /mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents/ --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 100
'''
args = parser.parse_args()
world_size = int(os.environ.get("WORLD_SIZE", 1))
maybe_init_distributed_environment_and_model_parallel(
tp_size=(world_size + 1) // 2, sp_size=(world_size + 1) // 2)
logger.info("Initialized distributed environment with world_size=%d",
world_size)
# Create DataLoader with proper settings
dataset, dataloader = build_parquet_iterable_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
logger.info("Initialized dataloader")
if args.verify_resume:
# First pass - record latent sums
first_pass_sums = []
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
latent_sum = latents.sum().item()
first_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f", i, latent_sum)
if i >= args.num_batches_per_epoch - 1:
break
# Save dataloader state using distributed checkpoint
checkpoint_dir = pathlib.Path(args.checkpoint_path)
logger.info("Rank %d: Saving dataloader state to %s", get_world_rank(),
checkpoint_dir)
states = {"dataloader": dataloader}
begin_time = time.monotonic()
dist_cp.save(states, checkpoint_id=checkpoint_dir.as_posix())
end_time = time.monotonic()
logger.info("Rank %d: Saved checkpoint in %.2f seconds",
get_world_rank(), end_time - begin_time)
# Make sure all processes wait for checkpoint to be saved
if world_size > 1:
dist.barrier()
# Recreate dataloader and load state
dataset, dataloader = build_parquet_iterable_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
load_states = {"dataloader": dataloader}
dist_cp.load(load_states, checkpoint_id=checkpoint_dir.as_posix())
logger.info("Rank %d: Loaded dataloader state from %s",
get_world_rank(), checkpoint_dir)
# Second pass - verify latent sums match
for i, (latents, embeddings, masks) in enumerate(dataloader):
latent_sum = latents.sum().item()
first_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f",
i + args.num_batches_per_epoch, latent_sum)
if i >= args.num_batches_per_epoch - 1:
break
dataset, dataloader = build_parquet_iterable_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
# Second pass - verify latent sums match
second_pass_sums = []
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
latent_sum = latents.sum().item()
second_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f (should match first pass: %f)",
i, latent_sum, first_pass_sums[i])
if i >= args.num_batches_per_epoch * 2 - 1:
break
# Verify all sums match
if all(
abs(a - b) < 1e-6
for a, b in zip(first_pass_sums, second_pass_sums)):
logger.info(
"All latent sums match between passes - resume verification successful!"
)
else:
raise ValueError(
"Latent sums do not match between passes - resume verification failed!"
)
start_time = time.time()
total_samples = 0
total_batches = 0
for _ in range(args.num_epoch):
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
if i >= args.num_batches_per_epoch:
break
# Move data to device
latents = latents.to(get_torch_device())
embeddings = embeddings.to(get_torch_device())
# Calculate actual batch size
batch_size = latents.size(0)
total_samples += batch_size
total_batches += 1
# Print progress only from rank 0
if get_world_rank() == 0 and (i + 1) % 10 == 0:
elapsed = time.time() - start_time
samples_per_sec = total_samples / elapsed
logger.info("Batch %d/%d, Speed: %.2f samples/sec", i + 1,
args.num_batches_per_epoch, samples_per_sec)
# Final statistics
if world_size > 1:
dist.barrier()
if get_world_rank() == 0:
elapsed = time.time() - start_time
samples_per_sec = total_samples / elapsed
logger.info("\nBenchmark Results:")
logger.info("Total time: %.2f seconds", elapsed)
logger.info("Total samples: %d", total_samples)
logger.info("Average speed: %.2f samples/sec", samples_per_sec)
logger.info("Time per batch: %.2f ms", elapsed / total_batches * 1000)
if __name__ == "__main__":
try:
main()
finally:
cleanup_dist_env_and_memory()
@@ -0,0 +1,187 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import os
import pathlib
import time
import torch
import torch.distributed as dist
import torch.distributed.checkpoint as dist_cp
from fastvideo.v1.dataset.parquet_dataset_map_style import (
build_parquet_map_style_dataloader)
from fastvideo.v1.distributed import get_world_rank
from fastvideo.v1.distributed.parallel_state import (
cleanup_dist_env_and_memory, get_torch_device,
maybe_init_distributed_environment_and_model_parallel)
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
def main() -> None:
torch.multiprocessing.set_start_method("spawn", force=True)
parser = argparse.ArgumentParser(
description="Benchmark parquet map style dataset loading speed")
parser.add_argument(
"--path",
type=str,
help="Path to parquet dataset",
)
parser.add_argument("--batch_size",
type=int,
default=4,
help="Batch size for DataLoader")
parser.add_argument("--num_data_workers",
type=int,
help="Number of DataLoader workers")
parser.add_argument("--num_epoch",
type=int,
default=2,
help="Number of epoches to benchmark")
parser.add_argument("--verify_resume",
action="store_true",
help="Verify resume")
parser.add_argument(
"--num_batches_per_epoch",
type=int,
default=1000,
help="Number of batches to benchmark",
)
parser.add_argument('--checkpoint_path',
type=str,
default='dataloader_checkpoint',
help='Path to save/load checkpoint')
'''
example launch command:
torchrun --nproc_per_node=1 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 4 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 3 --verify_resume
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path data/crush-smol/latents/combined_parquet_dataset --batch_size 2 --num_data_workers 1 --num_epoch 2 --num_batches_per_epoch 5 --verify_resume
torchrun --nproc_per_node=8 --master_port=12358 fastvideo/v1/dataset/benchmarks/benchmark_parquet_dataset_map_style.py --path /mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn/latents/ --batch_size 2 --num_data_workers 4 --num_epoch 2 --num_batches_per_epoch 100
'''
args = parser.parse_args()
world_size = int(os.environ.get("WORLD_SIZE", 1))
maybe_init_distributed_environment_and_model_parallel(
tp_size=(world_size + 1) // 2, sp_size=(world_size + 1) // 2)
logger.info("Initialized distributed environment with world_size=%d",
world_size)
# Create DataLoader with proper settings
dataset, dataloader = build_parquet_map_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
logger.info("Initialized dataloader with %d batches", len(dataloader))
if args.verify_resume:
# First pass - record latent sums
first_pass_sums = []
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
latent_sum = latents.sum().item()
first_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f", i, latent_sum)
if i >= args.num_batches_per_epoch - 1:
break
# Save dataloader state using distributed checkpoint
checkpoint_dir = pathlib.Path(args.checkpoint_path)
logger.info("Rank %d: Saving dataloader state to %s", get_world_rank(),
checkpoint_dir)
states = {"dataloader": dataloader}
begin_time = time.monotonic()
dist_cp.save(states, checkpoint_id=checkpoint_dir.as_posix())
end_time = time.monotonic()
logger.info("Rank %d: Saved checkpoint in %.2f seconds",
get_world_rank(), end_time - begin_time)
# Make sure all processes wait for checkpoint to be saved
if world_size > 1:
dist.barrier()
# Recreate dataloader and load state
dataset, dataloader = build_parquet_map_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
load_states = {"dataloader": dataloader}
dist_cp.load(load_states, checkpoint_id=checkpoint_dir.as_posix())
logger.info("Rank %d: Loaded dataloader state from %s",
get_world_rank(), checkpoint_dir)
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
latent_sum = latents.sum().item()
first_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f",
i + args.num_batches_per_epoch, latent_sum)
if i >= args.num_batches_per_epoch - 1:
break
dataset, dataloader = build_parquet_map_style_dataloader(
args.path, args.batch_size, args.num_data_workers)
# Second pass - verify latent sums match
second_pass_sums = []
for i, (latents, embeddings, masks) in enumerate(dataloader):
latent_sum = latents.sum().item()
second_pass_sums.append(latent_sum)
logger.info("Batch %d latent sum: %f (should match first pass: %f)",
i, latent_sum, first_pass_sums[i])
if i >= args.num_batches_per_epoch * 2 - 1:
break
# Verify all sums match
if all(
abs(a - b) < 1e-6
for a, b in zip(first_pass_sums, second_pass_sums)):
logger.info(
"All latent sums match between passes - resume verification successful!"
)
else:
raise ValueError(
"Latent sums do not match between passes - resume verification failed!"
)
start_time = time.time()
total_samples = 0
total_batches = 0
for _ in range(args.num_epoch):
for i, (latents, embeddings, masks,
caption_text) in enumerate(dataloader):
if i >= args.num_batches_per_epoch:
break
# Move data to device
latents = latents.to(get_torch_device())
embeddings = embeddings.to(get_torch_device())
# Calculate actual batch size
batch_size = latents.size(0)
total_samples += batch_size
total_batches += 1
# Print progress only from rank 0
if get_world_rank() == 0 and (i + 1) % 10 == 0:
elapsed = time.time() - start_time
samples_per_sec = total_samples / elapsed
logger.info("Batch %d/%d, Speed: %.2f samples/sec", i + 1,
args.num_batches_per_epoch, samples_per_sec)
# Final statistics
if world_size > 1:
dist.barrier()
if get_world_rank() == 0:
elapsed = time.time() - start_time
samples_per_sec = total_samples / elapsed
logger.info("\nBenchmark Results:")
logger.info("Total time: %.2f seconds", elapsed)
logger.info("Total samples: %d", total_samples)
logger.info("Average speed: %.2f samples/sec", samples_per_sec)
logger.info("Time per batch: %.2f ms", elapsed / total_batches * 1000)
if __name__ == "__main__":
try:
main()
finally:
cleanup_dist_env_and_memory()
-137
View File
@@ -1,137 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import json
import os
import time
from multiprocessing import Pool, cpu_count
from pathlib import Path
import torchvision
from tqdm import tqdm
def get_video_info(video_path):
"""Get video information using torchvision."""
# Read video tensor (T, C, H, W)
video_tensor, _, info = torchvision.io.read_video(str(video_path),
output_format="TCHW",
pts_unit="sec")
num_frames = video_tensor.shape[0]
height = video_tensor.shape[2]
width = video_tensor.shape[3]
fps = info.get("video_fps", 0)
duration = num_frames / fps if fps > 0 else 0
# Extract name
_, _, videos_dir, video_name = str(video_path).split("/")
return {
"path": str(video_name),
"resolution": {
"width": width,
"height": height
},
"size": os.path.getsize(video_path),
"fps": fps,
"duration": duration,
"num_frames": num_frames
}
def prepare_dataset_json(folder_path,
output_name="videos2caption.json",
num_workers=None) -> None:
"""Prepare dataset information from a folder containing videos and prompt.txt."""
folder_path = Path(folder_path)
# Read prompt file
prompt_file = folder_path / "prompt.txt"
if not prompt_file.exists():
raise FileNotFoundError(f"prompt.txt not found in {folder_path}")
with open(prompt_file) as f:
prompts = [line.strip() for line in f.readlines() if line.strip()]
# Read videos file
videos_file = folder_path / "videos.txt"
if not videos_file.exists():
raise FileNotFoundError(f"videos.txt not found in {folder_path}")
with open(videos_file) as f:
video_paths = [line.strip() for line in f.readlines() if line.strip()]
if len(prompts) != len(video_paths):
raise ValueError(
f"Number of prompts ({len(prompts)}) does not match number of videos ({len(video_paths)})"
)
# Prepare arguments for multiprocessing
process_args = [folder_path / video_path for video_path in video_paths]
# Determine number of workers
if num_workers is None:
num_workers = max(1, cpu_count() - 1) # Leave one CPU free
# Process videos in parallel
start_time = time.time()
with Pool(num_workers) as pool:
results = list(
tqdm(pool.imap(get_video_info, process_args),
total=len(process_args),
desc="Processing videos",
unit="video"))
# Combine results with prompts
dataset_info = []
for result, prompt in zip(results, prompts):
result["cap"] = [prompt]
dataset_info.append(result)
# Calculate total processing time
total_time = time.time() - start_time
total_videos = len(dataset_info)
avg_time_per_video = total_time / total_videos if total_videos > 0 else 0
print("\nProcessing completed:")
print(f"Total videos processed: {total_videos}")
print(f"Total time: {total_time:.2f} seconds")
print(f"Average time per video: {avg_time_per_video:.2f} seconds")
# Save to JSON file
output_file = folder_path / output_name
with open(output_file, 'w') as f:
json.dump(dataset_info, f, indent=2)
# Create merge.txt
merge_file = folder_path / "merge.txt"
with open(merge_file, 'w') as f:
f.write(f"{folder_path}/videos,{output_file}\n")
print(f"Dataset information saved to {output_file}")
print(f"Merge file created at {merge_file}")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description='Prepare video dataset information in JSON format')
parser.add_argument(
'--folder',
type=str,
required=True,
help='Path to the folder containing videos and prompt.txt')
parser.add_argument(
'--output',
type=str,
default='videos2caption.json',
help='Name of the output JSON file (default: videos2caption.json)')
parser.add_argument('--workers',
type=int,
default=32,
help='Number of worker processes (default: 16)')
return parser.parse_args()
if __name__ == "__main__":
args = parse_args()
prepare_dataset_json(args.folder, args.output, args.workers)
@@ -0,0 +1,275 @@
import os
import pickle
import random
from typing import Dict, List, Tuple
import numpy as np
import pyarrow.parquet as pq
import torch
import tqdm
from torch.utils.data import IterableDataset, get_worker_info
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.dataset.utils import collate_latents_embs_masks
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
get_world_size)
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class BatchIterator:
# TODO: Implement state_dict and load_state_dict to support resume.
def __init__(self, files, batch_size, text_padding_length, keys,
worker_num_samples, read_batch_size):
self.files = files
self.batch_size = batch_size
self.text_padding_length = text_padding_length
self.keys = keys
self.worker_num_samples = worker_num_samples
self.processed_samples = 0
self.buffer = []
self.read_batch_size = read_batch_size
def __iter__(self):
for file in self.files:
if self.processed_samples >= self.worker_num_samples:
return
reader = pq.ParquetFile(file)
for batch in reader.iter_batches(batch_size=self.read_batch_size):
if self.processed_samples >= self.worker_num_samples:
return
self.buffer.extend(batch.to_pylist())
while len(self.buffer) >= self.batch_size:
if self.processed_samples >= self.worker_num_samples:
return
batch_to_process = self.buffer[:self.batch_size]
self.buffer = self.buffer[self.batch_size:]
all_latents, all_embs, all_masks, caption_text = collate_latents_embs_masks(
batch_to_process, self.text_padding_length, self.keys)
self.processed_samples += self.batch_size
yield all_latents, all_embs, all_masks, caption_text
class LatentsParquetIterStyleDataset(IterableDataset):
"""Efficient loader for video-text data from a directory of Parquet files."""
# Modify this in the future if we want to add more keys, for example, in image to video.
keys = [("vae_latent", "latent"), ("text_embedding")]
def __init__(self,
path: str,
batch_size: int = 1024,
cfg_rate: float = 0.1,
num_workers: int = 1,
drop_last: bool = True,
text_padding_length: int = 512,
seed: int = 42,
read_batch_size: int = 32):
super().__init__()
self.path = str(path)
self.batch_size = batch_size
self.cfg_rate = cfg_rate
self.text_padding_length = text_padding_length
self.seed = seed
self.read_batch_size = read_batch_size
# Get distributed training info
self.global_rank = get_world_rank()
self.world_size = get_world_size()
self.sp_world_size = get_sp_world_size()
self.num_sp_groups = self.world_size // self.sp_world_size
num_workers = 1 if num_workers == 0 else num_workers
# Get sharding info
shard_parquet_files, shard_total_samples, shard_parquet_lengths = shard_parquet_files_across_sp_groups_and_workers(
self.path, self.num_sp_groups, num_workers, seed)
if drop_last:
self.worker_num_samples = min(
shard_total_samples) // batch_size * batch_size
# Assign files to current rank's SP group
ith_sp_group = self.global_rank // self.sp_world_size
self.sp_group_parquet_files = shard_parquet_files[ith_sp_group::self
.num_sp_groups]
self.sp_group_parquet_lengths = shard_parquet_lengths[
ith_sp_group::self.num_sp_groups]
self.sp_group_num_samples = shard_total_samples[ith_sp_group::self.
num_sp_groups]
logger.info(
"In total %d parquet files, %d samples, after sharding we retain %d samples due to drop_last",
sum([len(shard) for shard in shard_parquet_files]),
sum(shard_total_samples),
self.worker_num_samples * self.num_sp_groups * num_workers)
else:
raise ValueError("drop_last must be True")
logger.info("Each dataloader worker will load %d samples",
self.worker_num_samples)
def __iter__(self):
worker_info = get_worker_info()
worker_id = worker_info.id if worker_info is not None else 1
worker_files = self.sp_group_parquet_files[worker_id]
batch_iterator = BatchIterator(
files=worker_files,
batch_size=self.batch_size,
text_padding_length=self.text_padding_length,
keys=self.keys,
worker_num_samples=self.worker_num_samples,
read_batch_size=self.read_batch_size) # type: ignore
yield from batch_iterator
if batch_iterator.processed_samples != self.worker_num_samples:
raise ValueError(
"Rank %d, Worker %d: Not enough samples to process, this should not happen",
self.global_rank, worker_id)
def shard_parquet_files_across_sp_groups_and_workers(
path: str,
num_sp_groups: int,
num_workers: int,
seed: int = 42,
) -> Tuple[List[List[str]], List[int], List[Dict[str, int]]]:
"""
Shard parquet files across SP groups and workers in a balanced way.
Args:
path: Directory containing parquet files
num_sp_groups: Number of SP groups to shard across
num_workers: Number of workers per SP group
seed: Random seed for shuffling
Returns:
Tuple containing:
- List of lists of parquet files for each shard
- List of total samples per shard
- List of dictionaries mapping file paths to their lengths
"""
# Check if sharding plan already exists
sharding_info_dir = os.path.join(
path, f"sharding_info_{num_sp_groups}_sp_groups_{num_workers}_workers")
if os.path.exists(sharding_info_dir):
logger.info("Sharding plan already exists")
logger.info("Loading sharding plan from %s", sharding_info_dir)
try:
with open(
os.path.join(sharding_info_dir, "shard_parquet_files.pkl"),
"rb") as f:
shard_parquet_files = pickle.load(f)
with open(
os.path.join(sharding_info_dir, "shard_total_samples.pkl"),
"rb") as f:
shard_total_samples = pickle.load(f)
with open(
os.path.join(sharding_info_dir,
"shard_parquet_lengths.pkl"), "rb") as f:
shard_parquet_lengths = pickle.load(f)
return shard_parquet_files, shard_total_samples, shard_parquet_lengths
except Exception as e:
logger.error("Error loading sharding plan: %s", str(e))
logger.info("Falling back to creating new sharding plan")
if get_world_rank() == 0:
logger.info("Scanning for parquet files in %s", path)
# Find all parquet files
parquet_files = []
for root, _, files in os.walk(path):
for file in files:
if file.endswith('.parquet'):
parquet_files.append(os.path.join(root, file))
if not parquet_files:
raise ValueError("No parquet files found in %s", path)
# Calculate file lengths efficiently using a single pass
logger.info("Calculating file lengths...")
lengths = []
for file in tqdm.tqdm(parquet_files, desc="Reading parquet files"):
lengths.append(pq.ParquetFile(file).metadata.num_rows)
total_samples = sum(lengths)
logger.info("Found %d files with %d total samples", len(parquet_files),
total_samples)
# Sort files by length for better balancing
sorted_indices = np.argsort(lengths)
sorted_files = [parquet_files[i] for i in sorted_indices]
sorted_lengths = [lengths[i] for i in sorted_indices]
# Create shards
num_shards = num_sp_groups * num_workers
shard_parquet_files = [[] for _ in range(num_shards)]
shard_total_samples = [0] * num_shards
shard_parquet_lengths = [{} for _ in range(num_shards)]
# Distribute files to shards using a greedy approach
logger.info("Distributing files to shards...")
for file, length in zip(reversed(sorted_files),
reversed(sorted_lengths)):
# Find shard with minimum current length
target_shard = np.argmin(shard_total_samples)
shard_parquet_files[target_shard].append(file)
shard_total_samples[target_shard] += length
shard_parquet_lengths[target_shard][file] = length
#randomize each shard
for shard in shard_parquet_files:
random.seed(seed)
random.shuffle(shard)
save_dir = os.path.join(
path,
f"sharding_info_{num_sp_groups}_sp_groups_{num_workers}_workers")
os.makedirs(save_dir, exist_ok=True)
with open(os.path.join(save_dir, "shard_parquet_files.pkl"), "wb") as f:
pickle.dump(shard_parquet_files, f)
with open(os.path.join(save_dir, "shard_total_samples.pkl"), "wb") as f:
pickle.dump(shard_total_samples, f)
with open(os.path.join(save_dir, "shard_parquet_lengths.pkl"),
"wb") as f:
pickle.dump(shard_parquet_lengths, f)
logger.info("Saved sharding info to %s", save_dir)
# wait for all ranks to finish
torch.distributed.barrier()
# recursive call
return shard_parquet_files_across_sp_groups_and_workers(
path, num_sp_groups, num_workers, seed)
def build_parquet_iterable_style_dataloader(
path: str,
batch_size: int,
num_data_workers: int,
cfg_rate: float = 0.0,
drop_last: bool = True,
text_padding_length: int = 512,
seed: int = 42,
read_batch_size: int = 32
) -> Tuple[LatentsParquetIterStyleDataset, StatefulDataLoader]:
"""Build a dataloader for the LatentsParquetIterStyleDataset."""
dataset = LatentsParquetIterStyleDataset(
path=path,
batch_size=batch_size,
cfg_rate=cfg_rate,
num_workers=num_data_workers,
drop_last=drop_last,
text_padding_length=text_padding_length,
seed=seed,
read_batch_size=read_batch_size)
loader = StatefulDataLoader(
dataset,
batch_size=1,
num_workers=num_data_workers,
pin_memory=True,
)
return dataset, loader
@@ -0,0 +1,311 @@
# SPDX-License-Identifier: Apache-2.0
import os
import pickle
from typing import Any, Dict, List, Tuple
import pyarrow.parquet as pq
# Torch in general
import torch
import tqdm
# Dataset
from torch.utils.data import Dataset, Sampler
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.dataset.utils import collate_latents_embs_masks
from fastvideo.v1.distributed import (get_sp_world_size, get_world_rank,
get_world_size)
from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class DP_SP_BatchSampler(Sampler[List[int]]):
"""
A simple sequential batch sampler that yields batches of indices.
"""
def __init__(
self,
batch_size: int,
dataset_size: int,
num_sp_groups: int,
sp_world_size: int,
global_rank: int,
drop_last: bool = True,
drop_first_row: bool = False,
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 drop_first_row:
# drop 0 in global_indices
global_indices = global_indices[global_indices != 0]
self.dataset_size = self.dataset_size - 1
if self.drop_last:
# For drop_last=True, we:
# 1. Ensure total samples is divisible by (batch_size * num_sp_groups)
# 2. This guarantees each SP group gets same number of complete batches
# 3. Prevents uneven batch sizes across SP groups at end of epoch
num_batches = self.dataset_size // self.batch_size
num_global_batches = num_batches // self.num_sp_groups
global_indices = global_indices[:num_global_batches *
self.num_sp_groups *
self.batch_size]
else:
if self.dataset_size % (self.num_sp_groups * self.batch_size) != 0:
# add more indices to make it divisible by (batch_size * num_sp_groups)
padding_size = self.num_sp_groups * self.batch_size - (
self.dataset_size % (self.num_sp_groups * self.batch_size))
logger.info("Padding the dataset from %d to %d",
self.dataset_size, self.dataset_size + padding_size)
global_indices = torch.cat(
[global_indices, global_indices[:padding_size]])
# shard the indices to each sp group
ith_sp_group = self.global_rank // self.sp_world_size
sp_group_local_indices = global_indices[ith_sp_group::self.
num_sp_groups]
self.sp_group_local_indices = sp_group_local_indices
logger.info("Dataset size for each sp group: %d",
len(sp_group_local_indices))
def __iter__(self):
indices = self.sp_group_local_indices
for i in range(0, len(indices), self.batch_size):
batch_indices = indices[i:i + self.batch_size]
yield batch_indices.tolist()
def __len__(self):
return len(self.sp_group_local_indices) // self.batch_size
def get_parquet_files_and_length(path: str):
# Check if cached info exists
cache_dir = os.path.join(path, "map_style_cache")
cache_file = os.path.join(cache_dir, "file_info.pkl")
if os.path.exists(cache_file):
logger.info("Loading cached file info from %s", cache_file)
try:
with open(cache_file, "rb") as f:
file_names_sorted, lengths_sorted = pickle.load(f)
return file_names_sorted, lengths_sorted
except Exception as e:
logger.error("Error loading cached file info: %s", str(e))
logger.info("Falling back to scanning files")
# If no cache exists or loading failed, scan files
if get_world_rank() == 0:
lengths = []
file_names = []
for root, _, files in os.walk(path):
for file in sorted(files):
if file.endswith('.parquet'):
file_path = os.path.join(root, file)
file_names.append(file_path)
for file_path in tqdm.tqdm(file_names,
desc="Reading parquet files to get lengths"):
num_rows = pq.ParquetFile(file_path).metadata.num_rows
lengths.append(num_rows)
# sort according to file name to ensure all rank has the same order (in case os.walk is not sorted)
file_names_sorted, lengths_sorted = zip(
*sorted(zip(file_names, lengths), key=lambda x: x[0]))
assert len(
file_names_sorted) != 0, "No parquet files found in the dataset"
os.makedirs(cache_dir, exist_ok=True)
with open(cache_file, "wb") as f:
pickle.dump((file_names_sorted, lengths_sorted), f)
logger.info("Saved file info to %s", cache_file)
# Wait for rank 0 to finish saving
if get_world_size() > 1:
torch.distributed.barrier()
return get_parquet_files_and_length(path)
def read_row_from_parquet_file(parquet_files: List[str], global_row_idx: int,
lengths: List[int]) -> Dict[str, Any]:
'''
Read a row from a parquet file.
Args:
parquet_files: List[str]
global_row_idx: int
lengths: List[int]
Returns:
'''
# find the parquet file and local row index
cumulative = 0
for file_index in range(len(lengths)):
if cumulative + lengths[file_index] > global_row_idx:
local_row_idx = global_row_idx - cumulative
break
cumulative += lengths[file_index]
parquet_file = pq.ParquetFile(parquet_files[file_index])
# Calculate the row group to read into memory and the local idx
# This way we can avoid reading in the entire parquet file
cumulative = 0
for i in range(parquet_file.num_row_groups):
num_rows = parquet_file.metadata.row_group(i).num_rows
if cumulative + num_rows > local_row_idx:
row_group_index = i
local_index = local_row_idx - cumulative
break
cumulative += num_rows
row_group = parquet_file.read_row_group(row_group_index).to_pydict()
row_dict = {k: v[local_index] for k, v in row_group.items()}
del row_group
return row_dict
# ────────────────────────────────────────────────────────────────────────────
# 2. Dataset with batched __getitems__
# ────────────────────────────────────────────────────────────────────────────
class LatentsParquetMapStyleDataset(Dataset):
"""
Return latents[B,C,T,H,W] and embeddings[B,L,D] in pinned CPU memory.
Note:
Using parquet for map style dataset is not efficient, we mainly keep it for backward compatibility and debugging.
"""
# Modify this in the future if we want to add more keys, for example, in image to video.
keys = [("vae_latent", "latent"), "text_embedding"]
def __init__(
self,
path: str,
batch_size: int,
cfg_rate: float = 0.0,
seed: int = 42,
drop_last: bool = True,
drop_first_row: bool = False,
text_padding_length: int = 512,
):
super().__init__()
self.path = path
self.cfg_rate = cfg_rate
if cfg_rate > 0.0:
raise ValueError(
"cfg_rate > 0.0 is not supported for now because it will trigger bug when num_data_workers > 0"
)
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,
drop_first_row=drop_first_row,
seed=seed,
)
logger.info("Dataset initialized with %d parquet files and %d rows",
len(self.parquet_files), sum(self.lengths))
def get_validation_negative_prompt(
self) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, str]:
"""
Get the negative prompt for validation.
This method ensures the negative prompt is loaded and cached properly.
Returns the processed negative prompt data (latents, embeddings, masks, info).
"""
# Read first row from first parquet file
file_path = self.parquet_files[0]
row_idx = 0
# Read the negative prompt data
row_dict = read_row_from_parquet_file([file_path], row_idx,
[self.lengths[0]])
all_latents_list, all_embs_list, all_masks_list, caption_text_list = collate_latents_embs_masks(
[row_dict], self.text_padding_length, self.keys)
all_latents, all_embs, all_masks, caption_text = all_latents_list[
0], all_embs_list[0], all_masks_list[0], caption_text_list[0]
# add batch dimension
if len(all_embs.shape) == 2:
all_embs = all_embs.unsqueeze(0)
if len(all_masks.shape) == 1:
all_masks = all_masks.unsqueeze(0).unsqueeze(0)
return all_latents, all_embs, all_masks, caption_text
# PyTorch calls this ONLY because the batch_sampler yields a list
def __getitems__(self, indices: List[int]):
"""
Batch fetch using read_row_from_parquet_file for each index.
"""
rows = [
read_row_from_parquet_file(self.parquet_files, idx, self.lengths)
for idx in indices
]
all_latents, all_embs, all_masks, caption_text = collate_latents_embs_masks(
rows, self.text_padding_length, self.keys)
return all_latents, all_embs, all_masks, caption_text
def __len__(self):
return sum(self.lengths)
# ────────────────────────────────────────────────────────────────────────────
# 3. Loader helper – everything else stays just like your original trainer
# ────────────────────────────────────────────────────────────────────────────
def passthrough(batch):
return batch
def build_parquet_map_style_dataloader(
path,
batch_size,
num_data_workers,
cfg_rate=0.0,
drop_last=True,
drop_first_row=False,
text_padding_length=512,
seed=42) -> Tuple[LatentsParquetMapStyleDataset, StatefulDataLoader]:
dataset = LatentsParquetMapStyleDataset(
path,
batch_size,
cfg_rate=cfg_rate,
drop_last=drop_last,
drop_first_row=drop_first_row,
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
-467
View File
@@ -1,467 +0,0 @@
# SPDX-License-Identifier: Apache-2.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_sp_group, get_sp_parallel_rank,
get_sp_world_size, get_world_rank,
get_world_size)
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,
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.global_rank = get_world_rank()
self.rank_in_sp_group = get_sp_parallel_rank()
self.sp_group = get_sp_group()
self.sp_world_size = get_sp_world_size()
self.world_size = get_world_size()
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_world_size_{self.world_size}_sp_size_{self.sp_world_size}.json"
)
# group_ranks: a list of lists
# len(group_ranks) = self.world_size
# len(group_ranks[i]) = self.sp_world_size
# group_ranks[i] represents the ranks of the SP group for the i-th GPU
# For example, if self.world_size = 4, self.sp_world_size = 2, then
# group_ranks = [[0, 1], [0, 1], [2, 3], [2, 3]]
sp_group_ranks = get_sp_group().ranks
group_ranks: List[List] = [[] for _ in range(self.world_size)]
dist.all_gather_object(group_ranks, sp_group_ranks)
if self.global_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):
logger.info("Using existing plan from %s", self.plan_output_dir)
else:
logger.info("Creating new plan for %s", self.plan_output_dir)
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), (0, 1), (2, 3), (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.global_rank} (SP rank {self.rank_in_sp_group}): 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.rank_in_sp_group, :, :, :]
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.global_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.global_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.rank_in_sp_group, :, :, :]
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,
)
# 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")
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()
+85
View File
@@ -0,0 +1,85 @@
from typing import Any, Dict, List
import numpy as np
import torch
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
"""
Pad or crop an embedding [L, D] to exactly padding_length tokens.
Return:
- [L, D] tensor in pinned CPU memory
- [L] attention mask in pinned CPU memory
"""
L, D = t.shape
if padding_length > L: # pad
pad = torch.zeros(padding_length - L, D, dtype=t.dtype, device=t.device)
return torch.cat([t, pad], 0), torch.cat(
[torch.ones(L), torch.zeros(padding_length - L)], 0)
else: # crop
return t[:padding_length], torch.ones(padding_length)
def get_torch_tensors_from_row_dict(row_dict, keys) -> Dict[str, Any]:
"""
Get the latents and prompts from a row dictionary.
"""
return_dict = {}
for key in keys:
shape, bytes = None, None
if isinstance(key, tuple):
for k in key:
try:
shape = row_dict[f"{k}_shape"]
bytes = row_dict[f"{k}_bytes"]
except KeyError:
continue
key = key[0]
if shape is None or bytes is None:
raise ValueError(f"Key {key} not found in row_dict")
else:
shape = row_dict[f"{key}_shape"]
bytes = row_dict[f"{key}_bytes"]
# TODO (peiyuan): read precision
data = np.frombuffer(bytes, dtype=np.float32).reshape(shape).copy()
data = torch.from_numpy(data)
if len(data.shape) == 3:
B, L, D = data.shape
assert B == 1, "Batch size must be 1"
data = data.squeeze(0)
return_dict[key] = data
return return_dict
def collate_latents_embs_masks(
batch_to_process, text_padding_length,
keys) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, List[str]]:
# Initialize tensors to hold padded embeddings and masks
all_latents = []
all_embs = []
all_masks = []
caption_text = []
# Process each row individually
for i, row in enumerate(batch_to_process):
# Get tensors from row
data = get_torch_tensors_from_row_dict(row, keys)
latents, emb = data["vae_latent"], data["text_embedding"]
padded_emb, mask = pad(emb, text_padding_length)
# Store in batch tensors
all_latents.append(latents)
all_embs.append(padded_emb)
all_masks.append(mask)
# TODO(py): remove this once we fix preprocess
try:
caption_text.append(row["prompt"])
except KeyError:
caption_text.append(row["caption"])
# Pin memory for faster transfer to GPU
all_latents = torch.stack(all_latents)
all_embs = torch.stack(all_embs)
all_masks = torch.stack(all_masks)
return all_latents, all_embs, all_masks, caption_text
+6 -47
View File
@@ -4,9 +4,9 @@
import argparse
import dataclasses
import os
from typing import Any, Dict, List, Optional, cast
from typing import List, cast
from fastvideo import PipelineConfig, VideoGenerator
from fastvideo import VideoGenerator
from fastvideo.v1.configs.sample.base import SamplingParam
from fastvideo.v1.entrypoints.cli.cli_types import CLISubcommand
from fastvideo.v1.entrypoints.cli.utils import RaiseNotImplementedAction
@@ -37,8 +37,6 @@ class GenerateSubcommand(CLISubcommand):
def cmd(self, args: argparse.Namespace) -> None:
excluded_args = ['subparser', 'config', 'dispatch_function']
FastVideoArgs.from_cli_args(args)
provided_args = {}
for k, v in vars(args).items():
if (k not in excluded_args and v is not None
@@ -66,27 +64,19 @@ class GenerateSubcommand(CLISubcommand):
init_args = {
k: v
for k, v in merged_args.items() if k in self.init_arg_names
for k, v in merged_args.items()
if k not in self.generation_arg_names
}
generation_args = {
k: v
for k, v in merged_args.items() if k in self.generation_arg_names
}
pipeline_config = PipelineConfig.from_pretrained(
merged_args['model_path'])
update_config_from_args(pipeline_config.dit_config, merged_args,
"dit_config")
update_config_from_args(pipeline_config.vae_config, merged_args,
"vae_config")
update_config_from_args(pipeline_config, merged_args)
model_path = init_args.pop('model_path')
prompt = generation_args.pop('prompt')
generator = VideoGenerator.from_pretrained(
model_path=model_path, **init_args, pipeline_config=pipeline_config)
generator = VideoGenerator.from_pretrained(model_path=model_path,
**init_args)
generator.generate_video(prompt=prompt, **generation_args)
@@ -132,34 +122,3 @@ class GenerateSubcommand(CLISubcommand):
def cmd_init() -> List[CLISubcommand]:
return [GenerateSubcommand()]
def update_config_from_args(config: Any,
args_dict: Dict[str, Any],
prefix: Optional[str] = None) -> None:
"""
Update configuration object from arguments dictionary.
Args:
config: The configuration object to update
args_dict: Dictionary containing arguments
prefix: Prefix for the configuration parameters in the args_dict.
If None, assumes direct attribute mapping without prefix.
"""
# Handle top-level attributes (no prefix)
if prefix is None:
for key, value in args_dict.items():
if hasattr(config, key) and value is not None:
if key == "text_encoder_precisions" and isinstance(value, list):
setattr(config, key, tuple(value))
else:
setattr(config, key, value)
return
# Handle nested attributes with prefix
prefix_with_dot = f"{prefix}."
for key, value in args_dict.items():
if key.startswith(prefix_with_dot) and value is not None:
attr_name = key[len(prefix_with_dot):]
if hasattr(config, attr_name):
setattr(config, attr_name, value)
+12 -34
View File
@@ -18,8 +18,6 @@ import torch
import torchvision
from einops import rearrange
from fastvideo.v1.configs.pipelines import (PipelineConfig,
get_pipeline_config_cls_for_name)
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
@@ -55,9 +53,6 @@ class VideoGenerator:
model_path: str,
device: Optional[str] = None,
torch_dtype: Optional[torch.dtype] = None,
pipeline_config: Optional[
Union[str
| PipelineConfig]] = None,
**kwargs) -> "VideoGenerator":
"""
Create a video generator from a pretrained model.
@@ -66,35 +61,17 @@ class VideoGenerator:
model_path: Path or identifier for the pretrained model
device: Device to load the model on (e.g., "cuda", "cuda:0", "cpu")
torch_dtype: Data type for model weights (e.g., torch.float16)
**kwargs: Additional arguments to customize model loading
pipeline_config: Pipeline config to use for inference
**kwargs: Additional arguments to customize model loading, set any FastVideoArgs or PipelineConfig attributes here.
Returns:
The created video generator
Priority level: Default pipeline config < User's pipeline config < User's kwargs
"""
config = None
# 1. If users provide a pipeline config, it will override the default pipeline config
if isinstance(pipeline_config, PipelineConfig):
config = pipeline_config
else:
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
if isinstance(pipeline_config, str):
config.load_from_json(pipeline_config)
# 2. If users also provide some kwargs, it will override the pipeline config.
# The user kwargs shouldn't contain model config parameters!
if config is None:
logger.warning("No config found for model %s, using default config",
model_path)
config_args = kwargs
else:
config_args = shallow_asdict(config)
config_args.update(kwargs)
fastvideo_args = FastVideoArgs(model_path=model_path, **config_args)
# If users also provide some kwargs, it will override the FastVideoArgs and PipelineConfig.
kwargs['model_path'] = model_path
fastvideo_args = FastVideoArgs.from_kwargs(kwargs)
return cls.from_fastvideo_args(fastvideo_args)
@@ -150,16 +127,17 @@ class VideoGenerator:
"""
# Create a copy of inference args to avoid modifying the original
fastvideo_args = self.fastvideo_args
pipeline_config = fastvideo_args.pipeline_config
# Validate inputs
if not isinstance(prompt, str):
raise TypeError(
f"`prompt` must be a string, but got {type(prompt)}")
prompt = prompt.strip()
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
fastvideo_args.model_path)
kwargs["prompt"] = prompt
sampling_param.update(kwargs)
@@ -176,10 +154,10 @@ class VideoGenerator:
f"height={sampling_param.height}, width={sampling_param.width}, "
f"num_frames={sampling_param.num_frames}")
temporal_scale_factor = fastvideo_args.vae_config.arch_config.temporal_compression_ratio
temporal_scale_factor = pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = sampling_param.num_frames
num_gpus = fastvideo_args.num_gpus
use_temporal_scaling_frames = fastvideo_args.vae_config.use_temporal_scaling_frames
use_temporal_scaling_frames = pipeline_config.vae_config.use_temporal_scaling_frames
# Adjust number of frames based on number of GPUs
if use_temporal_scaling_frames:
@@ -238,18 +216,18 @@ class VideoGenerator:
num_videos_per_prompt: {sampling_param.num_videos_per_prompt}
guidance_scale: {sampling_param.guidance_scale}
n_tokens: {n_tokens}
flow_shift: {fastvideo_args.flow_shift}
embedded_guidance_scale: {fastvideo_args.embedded_cfg_scale}
flow_shift: {fastvideo_args.pipeline_config.flow_shift}
embedded_guidance_scale: {fastvideo_args.pipeline_config.embedded_cfg_scale}
save_video: {sampling_param.save_video}
output_path: {sampling_param.output_path}
""" # type: ignore[attr-defined]
logger.info(debug_str)
# Prepare batch
batch = ForwardBatch(
**shallow_asdict(sampling_param),
eta=0.0,
n_tokens=n_tokens,
VSA_sparsity=fastvideo_args.VSA_sparsity,
extra={},
)
+89 -226
View File
@@ -6,26 +6,32 @@ import argparse
import dataclasses
from contextlib import contextmanager
from dataclasses import field
from typing import Any, Callable, List, Optional, Tuple
from typing import Any, Dict, List, Optional
from fastvideo.v1.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.v1.configs.pipelines.base import PipelineConfig
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import FlexibleArgumentParser, StoreBoolean
logger = init_logger(__name__)
def preprocess_text(prompt: str) -> str:
return prompt
def clean_cli_args(args: argparse.Namespace) -> Dict[str, Any]:
"""
Clean the arguments by removing the ones that not explicitly provided by the user.
"""
provided_args = {}
for k, v in vars(args).items():
if (v is not None and hasattr(args, '_provided')
and k in args._provided):
provided_args[k] = v
def postprocess_text(output: Any) -> Any:
raise NotImplementedError
return provided_args
# args for fastvideo framework
@dataclasses.dataclass
class FastVideoArgs:
# Model and path configuration
# Model and path configuration (for convenience)
model_path: str
# Cache strategy
@@ -44,70 +50,29 @@ class FastVideoArgs:
num_gpus: int = 1
tp_size: int = -1
sp_size: int = -1
dp_size: int = 1
dp_shards: int = -1
hsdp_replicate_dim: int = 1
hsdp_shard_dim: int = -1
dist_timeout: Optional[int] = None # timeout for torch.distributed
# Video generation parameters
embedded_cfg_scale: float = 6.0
flow_shift: Optional[float] = None
pipeline_config: PipelineConfig = field(default_factory=PipelineConfig)
output_type: str = "pil"
# DiT configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
precision: str = "bf16"
use_cpu_offload: bool = True
use_fsdp_inference: bool = True
# VAE configuration
vae_precision: str = "fp16"
vae_tiling: bool = True # Might change in between forward passes
vae_sp: bool = False # Might change in between forward passes
# vae_scale_factor: Optional[int] = None # Deprecated
vae_config: VAEConfig = field(default_factory=VAEConfig)
# Image encoder configuration
image_encoder_precision: str = "fp32"
image_encoder_config: EncoderConfig = field(default_factory=EncoderConfig)
# Text encoder configuration
DEFAULT_TEXT_ENCODER_PRECISIONS = (
"fp16",
# "fp16",
)
text_encoder_precisions: Tuple[str, ...] = field(
default_factory=lambda: FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS)
text_encoder_configs: Tuple[EncoderConfig, ...] = field(
default_factory=lambda: (EncoderConfig(), ))
preprocess_text_funcs: Tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: Tuple[Callable[[Any], Any], ...] = field(
default_factory=lambda: (postprocess_text, ))
# STA parameters
# STA (Sliding Tile Attention) parameters
mask_strategy_file_path: Optional[str] = None
STA_mode: Optional[str] = None
skip_time_steps: int = 15
# LoRA parameters
lora_path: Optional[str] = None
lora_nickname: Optional[
str] = "default" # for swapping adapters in the pipeline
lora_target_names: Optional[List[
str]] = None # can restrict list of layers to adapt, e.g. ["q_proj"]
# STA parameters
mask_strategy_file_path: Optional[str] = None
# Compilation
enable_torch_compile: bool = False
disable_autocast: bool = False
# StepVideo specific parameters
pos_magic: Optional[str] = None
neg_magic: Optional[str] = None
timesteps_scale: Optional[bool] = None
# Logging
log_level: str = "info"
# VSA parameters
VSA_sparsity: float = 0.0 # inference/validation sparsity
@property
def training_mode(self) -> bool:
@@ -125,11 +90,6 @@ class FastVideoArgs:
help=
"The path of the model weights. This can be a local folder or a Hugging Face repo ID.",
)
parser.add_argument(
"--dit-weight",
type=str,
help="Path to the DiT model weights",
)
parser.add_argument(
"--model-dir",
type=str,
@@ -175,31 +135,27 @@ class FastVideoArgs:
help="The number of GPUs to use.",
)
parser.add_argument(
"--tensor-parallel-size",
"--tp-size",
type=int,
default=FastVideoArgs.tp_size,
help="The tensor parallelism size.",
)
parser.add_argument(
"--sequence-parallel-size",
"--sp-size",
type=int,
default=FastVideoArgs.sp_size,
help="The sequence parallelism size.",
)
parser.add_argument(
"--data-parallel-size",
"--dp-size",
"--hsdp-replicate-dim",
type=int,
default=FastVideoArgs.dp_size,
default=FastVideoArgs.hsdp_replicate_dim,
help="The data parallelism size.",
)
parser.add_argument(
"--data-parallel-shards",
"--dp-shards",
"--hsdp-shard-dim",
type=int,
default=FastVideoArgs.dp_shards,
default=FastVideoArgs.hsdp_shard_dim,
help="The data parallelism shards.",
)
parser.add_argument(
@@ -209,19 +165,7 @@ class FastVideoArgs:
help="Set timeout for torch.distributed initialization.",
)
parser.add_argument(
"--embedded-cfg-scale",
type=float,
default=FastVideoArgs.embedded_cfg_scale,
help="Embedded CFG scale",
)
parser.add_argument(
"--flow-shift",
"--shift",
type=float,
default=FastVideoArgs.flow_shift,
help="Flow shift parameter",
)
# Output type
parser.add_argument(
"--output-type",
type=str,
@@ -230,53 +174,7 @@ class FastVideoArgs:
help="Output type for the generated video",
)
parser.add_argument(
"--precision",
type=str,
default=FastVideoArgs.precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for the model",
)
# VAE configuration
parser.add_argument(
"--vae-precision",
type=str,
default=FastVideoArgs.vae_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for VAE",
)
parser.add_argument(
"--vae-tiling",
action=StoreBoolean,
default=FastVideoArgs.vae_tiling,
help="Enable VAE tiling",
)
parser.add_argument(
"--vae-sp",
action=StoreBoolean,
help="Enable VAE spatial parallelism",
)
parser.add_argument(
"--text-encoder-precisions",
nargs="+",
type=str,
default=FastVideoArgs.DEFAULT_TEXT_ENCODER_PRECISIONS,
choices=["fp32", "fp16", "bf16"],
help="Precision for each text encoder",
)
# Image encoder config
parser.add_argument(
"--image-encoder-precision",
type=str,
default=FastVideoArgs.image_encoder_precision,
choices=["fp32", "fp16", "bf16"],
help="Precision for image encoder",
)
# STA parameters
# STA (Sliding Tile Attention) parameters
parser.add_argument(
"--STA-mode",
type=str,
@@ -325,90 +223,61 @@ class FastVideoArgs:
"Disable autocast for denoising loop and vae decoding in pipeline sampling",
)
# VSA parameters
parser.add_argument(
"--pos_magic",
type=str,
default=FastVideoArgs.pos_magic,
help="Positive magic prompt for sampling",
)
parser.add_argument(
"--neg_magic",
type=str,
default=FastVideoArgs.neg_magic,
help="Negative magic prompt for sampling",
)
parser.add_argument(
"--timesteps_scale",
type=bool,
default=FastVideoArgs.timesteps_scale,
help="Bool for applying scheduler scale in set_timesteps",
"--VSA-sparsity",
type=float,
default=FastVideoArgs.VSA_sparsity,
help="Validation sparsity for VSA",
)
# Logging
parser.add_argument(
"--log-level",
type=str,
default=FastVideoArgs.log_level,
help="The logging level of all loggers.",
)
# Add VAE configuration arguments
from fastvideo.v1.configs.models.vaes.base import VAEConfig
VAEConfig.add_cli_args(parser)
# Add DiT configuration arguments
from fastvideo.v1.configs.models.dits.base import DiTConfig
DiTConfig.add_cli_args(parser)
# Add pipeline configuration arguments
PipelineConfig.add_cli_args(parser)
return parser
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "FastVideoArgs":
args.tp_size = args.tensor_parallel_size
args.sp_size = args.sequence_parallel_size
args.flow_shift = getattr(args, "shift", args.flow_shift)
provided_args = clean_cli_args(args)
# Get all fields from the dataclass
attrs = [attr.name for attr in dataclasses.fields(cls)]
# Create a dictionary of attribute values, with defaults for missing attributes
kwargs = {}
for attr in attrs:
# Handle renamed attributes or those with multiple CLI names
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
kwargs[attr] = args.tensor_parallel_size
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
kwargs[attr] = args.sequence_parallel_size
elif attr == 'dp_size' and hasattr(args, 'data_parallel_size'):
kwargs[attr] = args.data_parallel_size
elif attr == 'dp_shards' and hasattr(args, 'data_parallel_shards'):
kwargs[attr] = args.data_parallel_shards
elif attr == 'flow_shift' and hasattr(args, 'shift'):
kwargs[attr] = args.shift
if attr == 'pipeline_config':
pipeline_config = PipelineConfig.from_kwargs(provided_args)
kwargs[attr] = pipeline_config
# Use getattr with default value from the dataclass for potentially missing attributes
else:
default_value = getattr(cls, attr, None)
if getattr(args, attr, default_value) is not None:
kwargs[attr] = getattr(args, attr, default_value)
value = getattr(args, attr, default_value)
kwargs[attr] = value # type: ignore
return cls(**kwargs) # type: ignore
@classmethod
def from_kwargs(cls, kwargs: Dict[str, Any]) -> "FastVideoArgs":
kwargs['pipeline_config'] = PipelineConfig.from_kwargs(kwargs)
return cls(**kwargs)
def check_fastvideo_args(self) -> None:
"""Validate inference arguments for consistency"""
if not self.inference_mode:
assert self.dp_size is not -1, "dp_size must be set for training"
assert self.dp_shards is not -1, "dp_shards must be set for training"
assert self.sp_size is not -1, "sp_size must be set for training"
assert self.hsdp_replicate_dim != -1, "hsdp_replicate_dim must be set for training"
assert self.hsdp_shard_dim != -1, "hsdp_shard_dim must be set for training"
assert self.sp_size != -1, "sp_size must be set for training"
if self.tp_size is -1:
if self.tp_size == -1:
self.tp_size = self.num_gpus
if self.sp_size is -1:
if self.sp_size == -1:
self.sp_size = self.num_gpus
if self.dp_shards is -1:
self.dp_shards = self.num_gpus
if self.hsdp_shard_dim == -1:
self.hsdp_shard_dim = self.num_gpus
assert self.sp_size <= self.num_gpus and self.num_gpus % self.sp_size == 0, "num_gpus must >= and be divisible by sp_size"
assert self.dp_size <= self.num_gpus and self.num_gpus % self.dp_size == 0, "num_gpus must >= and be divisible by dp_size"
assert self.dp_shards <= self.num_gpus and self.num_gpus % self.dp_shards == 0, "num_gpus must >= and be divisible by dp_shards"
assert self.hsdp_replicate_dim <= self.num_gpus and self.num_gpus % self.hsdp_replicate_dim == 0, "num_gpus must >= and be divisible by hsdp_replicate_dim"
assert self.hsdp_shard_dim <= self.num_gpus and self.num_gpus % self.hsdp_shard_dim == 0, "num_gpus must >= and be divisible by hsdp_shard_dim"
if self.num_gpus < max(self.tp_size, self.sp_size):
self.num_gpus = max(self.tp_size, self.sp_size)
@@ -418,33 +287,17 @@ class FastVideoArgs:
f"tp_size ({self.tp_size}) must be equal to sp_size ({self.sp_size})"
)
# Validate VAE spatial parallelism with VAE tiling
if self.vae_sp and not self.vae_tiling:
raise ValueError(
"Currently enabling vae_sp requires enabling vae_tiling, please set --vae-tiling to True."
)
if len(self.text_encoder_configs) != len(self.text_encoder_precisions):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text encoder precisions ({len(self.text_encoder_precisions)})"
)
if len(self.text_encoder_configs) != len(self.preprocess_text_funcs):
raise ValueError(
f"Length of text encoder configs ({len(self.text_encoder_configs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
if len(self.preprocess_text_funcs) != len(self.postprocess_text_funcs):
raise ValueError(
f"Length of text postprocess functions ({len(self.postprocess_text_funcs)}) must be equal to length of text preprocessing functions ({len(self.preprocess_text_funcs)})"
)
if self.enable_torch_compile and self.num_gpus > 1:
logger.warning(
"Currently torch compile does not work with multi-gpu. Setting enable_torch_compile to False"
)
self.enable_torch_compile = False
if self.pipeline_config is None:
raise ValueError("pipeline_config is not set in FastVideoArgs")
self.pipeline_config.check_pipeline_config()
_current_fastvideo_args = None
@@ -518,7 +371,6 @@ class TrainingArgs(FastVideoArgs):
# text encoder & vae & diffusion model
pretrained_model_name_or_path: str = ""
dit_model_name_or_path: str = ""
cache_dir: str = ""
# diffusion setting
ema_decay: float = 0.0
@@ -533,6 +385,7 @@ class TrainingArgs(FastVideoArgs):
validation_steps: float = 0.0
log_validation: bool = False
tracker_project_name: str = ""
wandb_run_name: str = ""
seed: Optional[int] = None
# output
@@ -540,7 +393,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
@@ -581,34 +433,29 @@ class TrainingArgs(FastVideoArgs):
# master_weight_type
master_weight_type: str = ""
# For fast checking in LoRA pipeline
training_mode: bool = True
# VSA training decay parameters
VSA_decay_rate: float = 0.01 # decay rate -> 0.02
VSA_decay_interval_steps: int = 1 # decay interval steps -> 50
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
provided_args = clean_cli_args(args)
# Get all fields from the dataclass
attrs = [attr.name for attr in dataclasses.fields(cls)]
logger.info(provided_args)
# Create a dictionary of attribute values, with defaults for missing attributes
kwargs = {}
for attr in attrs:
# Handle renamed attributes or those with multiple CLI names
if attr == 'tp_size' and hasattr(args, 'tensor_parallel_size'):
kwargs[attr] = args.tensor_parallel_size
elif attr == 'sp_size' and hasattr(args, 'sequence_parallel_size'):
kwargs[attr] = args.sequence_parallel_size
elif attr == 'flow_shift' and hasattr(args, 'shift'):
kwargs[attr] = args.shift
elif attr == 'dp_size' and hasattr(args, 'data_parallel_size'):
kwargs[attr] = args.data_parallel_size
elif attr == 'dp_shards' and hasattr(args, 'data_parallel_shards'):
kwargs[attr] = args.data_parallel_shards
if attr == 'pipeline_config':
pipeline_config = PipelineConfig.from_kwargs(provided_args)
kwargs[attr] = pipeline_config
# Use getattr with default value from the dataclass for potentially missing attributes
else:
default_value = getattr(cls, attr, None)
kwargs[attr] = getattr(args, attr, default_value)
value = getattr(args, attr, default_value)
kwargs[attr] = value # type: ignore
return cls(**kwargs)
return cls(**kwargs) # type: ignore
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
@@ -698,8 +545,12 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--tracker-project-name",
type=str,
help="Project name for tracking")
parser.add_argument("--wandb-run-name",
type=str,
help="Run name for wandb")
parser.add_argument("--seed",
type=int,
default=42,
help="Seed for deterministic training")
# Output configuration
@@ -835,4 +686,16 @@ class TrainingArgs(FastVideoArgs):
type=str,
help="Master weight type")
# VSA parameters for training with dense to sparse adaption
parser.add_argument(
"--VSA-decay-rate", # decay rate, how much sparsity you want to decay each step
type=float,
default=TrainingArgs.VSA_decay_rate,
help="VSA decay rate")
parser.add_argument(
"--VSA-decay-interval-steps", # how many steps for training with current sparsity
type=int,
default=TrainingArgs.VSA_decay_interval_steps,
help="VSA decay interval steps")
return parser
+2 -2
View File
@@ -114,7 +114,7 @@ def _info(logger: Logger,
if (main_process_only and is_main_process) or (local_main_process_only
and is_local_main_process):
logger.log(logging.INFO, msg, *args, **kwargs)
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
global _warned_local_main_process, _warned_main_process
@@ -134,7 +134,7 @@ def _info(logger: Logger,
_warned_main_process = True
if not main_process_only and not local_main_process_only:
logger.log(logging.INFO, msg, *args, **kwargs)
logger.log(logging.INFO, msg, *args, stacklevel=2, **kwargs)
class _FastvideoLogger(Logger):
+4 -4
View File
@@ -6,7 +6,7 @@ import torch
from torch import nn
from fastvideo.v1.configs.models import DiTConfig
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
# TODO
@@ -19,7 +19,7 @@ class BaseDiT(nn.Module, ABC):
num_channels_latents: int
# always supports torch_sdpa
_supported_attention_backends: Tuple[
_Backend, ...] = DiTConfig()._supported_attention_backends
AttentionBackendEnum, ...] = DiTConfig()._supported_attention_backends
def __init_subclass__(cls) -> None:
required_class_attrs = [
@@ -65,7 +65,7 @@ class BaseDiT(nn.Module, ABC):
)
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
return self._supported_attention_backends
@@ -85,7 +85,7 @@ class CachableDiT(BaseDiT):
num_channels_latents: int
# always supports torch_sdpa
_supported_attention_backends: Tuple[
_Backend, ...] = DiTConfig()._supported_attention_backends
AttentionBackendEnum, ...] = DiTConfig()._supported_attention_backends
def __init__(self, config: DiTConfig, **kwargs) -> None:
super().__init__(config, **kwargs)
+7 -5
View File
@@ -23,7 +23,7 @@ from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
unpatchify)
from fastvideo.v1.models.dits.base import CachableDiT
from fastvideo.v1.models.utils import modulate
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
class HunyuanRMSNorm(nn.Module):
@@ -96,7 +96,8 @@ class MMDoubleStreamBlock(nn.Module):
num_attention_heads: int,
mlp_ratio: float,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
prefix: str = "",
):
super().__init__()
@@ -303,7 +304,8 @@ class MMSingleStreamBlock(nn.Module):
num_attention_heads: int,
mlp_ratio: float = 4.0,
dtype: Optional[torch.dtype] = None,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
prefix: str = "",
):
super().__init__()
@@ -876,8 +878,8 @@ class IndividualTokenRefinerBlock(nn.Module):
num_heads=num_attention_heads,
head_size=hidden_size // num_attention_heads,
# TODO: remove hardcode; remove STA
supported_attention_backends=(_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA),
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA),
)
def forward(self, x, c):
+14 -12
View File
@@ -26,7 +26,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.v1.layers.visual_embedding import TimestepEmbedder
from fastvideo.v1.models.dits.base import BaseDiT
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
class PatchEmbed2D(nn.Module):
@@ -139,16 +139,17 @@ class StepVideoRMSNorm(nn.Module):
class SelfAttention(nn.Module):
def __init__(self,
hidden_dim,
head_dim,
rope_split: Tuple[int, int, int] = (64, 32, 32),
bias: bool = False,
with_rope: bool = True,
with_qk_norm: bool = True,
attn_type: str = "torch",
supported_attention_backends=(_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA)):
def __init__(
self,
hidden_dim,
head_dim,
rope_split: Tuple[int, int, int] = (64, 32, 32),
bias: bool = False,
with_rope: bool = True,
with_qk_norm: bool = True,
attn_type: str = "torch",
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)):
super().__init__()
self.head_dim = head_dim
self.hidden_dim = hidden_dim
@@ -257,7 +258,8 @@ class CrossAttention(nn.Module):
head_dim,
bias=False,
with_qk_norm=True,
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)
) -> None:
super().__init__()
self.head_dim = head_dim
+29 -26
View File
@@ -26,7 +26,7 @@ from fastvideo.v1.layers.rotary_embedding import (_apply_rotary_emb,
from fastvideo.v1.layers.visual_embedding import (ModulateProjection,
PatchEmbed, TimestepEmbedder)
from fastvideo.v1.models.dits.base import CachableDiT
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
class WanImageEmbedding(torch.nn.Module):
@@ -125,8 +125,8 @@ class WanSelfAttention(nn.Module):
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=(_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA))
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA))
def forward(self, x: torch.Tensor, context: torch.Tensor,
context_lens: int):
@@ -174,7 +174,8 @@ class WanI2VCrossAttention(WanSelfAttention):
window_size=(-1, -1),
qk_norm=True,
eps=1e-6,
supported_attention_backends: Optional[Tuple[_Backend, ...]] = None
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None
) -> None:
super().__init__(dim, num_heads, window_size, qk_norm, eps,
supported_attention_backends)
@@ -216,17 +217,18 @@ class WanI2VCrossAttention(WanSelfAttention):
class WanTransformerBlock(nn.Module):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
prefix: str = ""):
def __init__(
self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
prefix: str = ""):
super().__init__()
# 1. Self-attention
@@ -358,17 +360,18 @@ class WanTransformerBlock(nn.Module):
class WanTransformerBlock_VSA(nn.Module):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[_Backend,
...]] = None,
prefix: str = ""):
def __init__(
self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: Optional[int] = None,
supported_attention_backends: Optional[Tuple[AttentionBackendEnum,
...]] = None,
prefix: str = ""):
super().__init__()
# 1. Self-attention
+7 -5
View File
@@ -8,12 +8,13 @@ from torch import nn
from fastvideo.v1.configs.models.encoders import (BaseEncoderOutput,
ImageEncoderConfig,
TextEncoderConfig)
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
class TextEncoder(nn.Module, ABC):
_supported_attention_backends: Tuple[
_Backend, ...] = TextEncoderConfig()._supported_attention_backends
AttentionBackendEnum,
...] = TextEncoderConfig()._supported_attention_backends
def __init__(self, config: TextEncoderConfig) -> None:
super().__init__()
@@ -34,13 +35,14 @@ class TextEncoder(nn.Module, ABC):
pass
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
return self._supported_attention_backends
class ImageEncoder(nn.Module, ABC):
_supported_attention_backends: Tuple[
_Backend, ...] = ImageEncoderConfig()._supported_attention_backends
AttentionBackendEnum,
...] = ImageEncoderConfig()._supported_attention_backends
def __init__(self, config: ImageEncoderConfig) -> None:
super().__init__()
@@ -56,5 +58,5 @@ class ImageEncoder(nn.Module, ABC):
pass
@property
def supported_attention_backends(self) -> Tuple[_Backend, ...]:
def supported_attention_backends(self) -> Tuple[AttentionBackendEnum, ...]:
return self._supported_attention_backends
+3 -5
View File
@@ -81,10 +81,7 @@ def get_hf_config(
return config
def get_diffusers_config(
model: str,
fastvideo_args: Optional[dict] = None,
) -> Dict[str, Any]:
def get_diffusers_config(model: str, ) -> Dict[str, Any]:
"""Gets a configuration for the given diffusers model.
Args:
@@ -105,7 +102,8 @@ def get_diffusers_config(
# Load the config directly from the file
with open(config_file) as f:
config_dict: Dict[str, Any] = json.load(f)
if "_diffusers_version" in config_dict:
config_dict.pop("_diffusers_version")
# TODO(will): apply any overrides from inference args
return config_dict
except Exception as e:
+38 -44
View File
@@ -15,8 +15,9 @@ from safetensors.torch import load_file as safetensors_load_file
from transformers import AutoImageProcessor, AutoTokenizer
from transformers.utils import SAFE_WEIGHTS_INDEX_NAME
from fastvideo.v1.configs.models import EncoderConfig
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.hf_transformer_utils import get_diffusers_config
from fastvideo.v1.models.loader.fsdp_load import maybe_load_fsdp_model
@@ -45,7 +46,7 @@ class ComponentLoader(ABC):
Args:
model_path: Path to the component model
architecture: Architecture of the component model
fastvideo_args: Inference arguments
fastvideo_args: FastVideoArgs
Returns:
The loaded component
@@ -183,9 +184,10 @@ class TextEncoderLoader(ComponentLoader):
self,
model_config: Any,
model: nn.Module,
model_path: str,
) -> Generator[Tuple[str, torch.Tensor], None, None]:
primary_weights = TextEncoderLoader.Source(
model_config.model,
model_path,
prefix="",
fall_back_to_pt=getattr(model, "fall_back_to_pt_during_load", True),
allow_patterns_overrides=getattr(model, "allow_patterns_overrides",
@@ -209,8 +211,7 @@ class TextEncoderLoader(ComponentLoader):
# revision=fastvideo_args.revision,
# model_override_args=None,
# )
with open(os.path.join(model_path, "config.json")) as f:
model_config = json.load(f)
model_config = get_diffusers_config(model=model_path)
model_config.pop("_name_or_path", None)
model_config.pop("transformers_version", None)
model_config.pop("model_type", None)
@@ -220,13 +221,17 @@ class TextEncoderLoader(ComponentLoader):
# @TODO(Wei): Better way to handle this?
try:
encoder_config = fastvideo_args.text_encoder_configs[0]
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[
0]
encoder_config.update_model_arch(model_config)
encoder_precision = fastvideo_args.text_encoder_precisions[0]
encoder_precision = fastvideo_args.pipeline_config.text_encoder_precisions[
0]
except Exception:
encoder_config = fastvideo_args.text_encoder_configs[1]
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[
1]
encoder_config.update_model_arch(model_config)
encoder_precision = fastvideo_args.text_encoder_precisions[1]
encoder_precision = fastvideo_args.pipeline_config.text_encoder_precisions[
1]
target_device = get_torch_device()
# TODO(will): add support for other dtypes
@@ -235,7 +240,7 @@ class TextEncoderLoader(ComponentLoader):
def load_model(self,
model_path: str,
model_config,
model_config: EncoderConfig,
target_device: torch.device,
dtype: str = "fp16"):
with set_default_torch_dtype(PRECISION_TO_TYPE[dtype]):
@@ -245,9 +250,8 @@ class TextEncoderLoader(ComponentLoader):
model = model_cls(model_config)
weights_to_load = {name for name, _ in model.named_parameters()}
model_config.model = model_path
loaded_weights = model.load_weights(
self._get_all_weights(model_config, model))
self._get_all_weights(model_config, model, model_path))
self.counter_after_loading_weights = time.perf_counter()
logger.info(
"Loading weights took %.2f seconds",
@@ -261,7 +265,6 @@ class TextEncoderLoader(ComponentLoader):
raise ValueError("Following weights were not initialized from "
f"checkpoint: {weights_not_loaded}")
# TODO(will): add support for training/finetune
return model.eval()
@@ -284,13 +287,14 @@ class ImageEncoderLoader(TextEncoderLoader):
model_config.pop("model_type", None)
logger.info("HF Model config: %s", model_config)
encoder_config = fastvideo_args.image_encoder_config
encoder_config = fastvideo_args.pipeline_config.image_encoder_config
encoder_config.update_model_arch(model_config)
target_device = get_torch_device()
# TODO(will): add support for other dtypes
return self.load_model(model_path, encoder_config, target_device,
fastvideo_args.image_encoder_precision)
return self.load_model(
model_path, encoder_config, target_device,
fastvideo_args.pipeline_config.image_encoder_precision)
class ImageProcessorLoader(ComponentLoader):
@@ -332,18 +336,17 @@ class VAELoader(ComponentLoader):
def load(self, model_path: str, architecture: str,
fastvideo_args: FastVideoArgs):
"""Load the VAE based on the model path, architecture, and inference args."""
# TODO(will): move this to a constants file
config = get_diffusers_config(model=model_path)
class_name = config.pop("_class_name")
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
config.pop("_diffusers_version")
vae_config = fastvideo_args.vae_config
vae_config = fastvideo_args.pipeline_config.vae_config
vae_config.update_model_arch(config)
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(get_torch_device())
with set_default_torch_dtype(PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]):
vae_cls, _ = ModelRegistry.resolve_model_cls(class_name)
vae = vae_cls(vae_config).to(get_torch_device())
# Find all safetensors files
safetensors_list = glob.glob(
@@ -355,10 +358,8 @@ class VAELoader(ComponentLoader):
loaded = safetensors_load_file(safetensors_list[0])
vae.load_state_dict(
loaded, strict=False) # We might only load encoder or decoder
dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
vae = vae.eval().to(dtype)
return vae
return vae.eval()
class TransformerLoader(ComponentLoader):
@@ -374,10 +375,9 @@ class TransformerLoader(ComponentLoader):
raise ValueError(
"Model config does not contain a _class_name attribute. "
"Only diffusers format is supported.")
config.pop("_diffusers_version")
# Config from Diffusers supersedes fastvideo's model config
dit_config = fastvideo_args.dit_config
dit_config = fastvideo_args.pipeline_config.dit_config
dit_config.update_model_arch(config)
model_cls, _ = ModelRegistry.resolve_model_cls(cls_name)
@@ -391,19 +391,13 @@ class TransformerLoader(ComponentLoader):
logger.info("Loading model from %s safetensors files in %s",
len(safetensors_list), model_path)
# initialize_sequence_parallel_group(fastvideo_args.sp_size)
if fastvideo_args.training_mode:
assert isinstance(
fastvideo_args, TrainingArgs
), "fastvideo_args must be a TrainingArgs object when training_mode is True"
default_dtype = PRECISION_TO_TYPE[fastvideo_args.master_weight_type]
else:
default_dtype = PRECISION_TO_TYPE[fastvideo_args.precision]
default_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.dit_precision]
# Load the model using FSDP loader
logger.info("Loading model from %s, default_dtype: %s", cls_name,
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={
@@ -412,8 +406,8 @@ class TransformerLoader(ComponentLoader):
},
weight_dir_list=safetensors_list,
device=get_torch_device(),
data_parallel_size=fastvideo_args.dp_size,
data_parallel_shards=fastvideo_args.dp_shards,
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,
@@ -462,15 +456,15 @@ class SchedulerLoader(ComponentLoader):
class_name = config.pop("_class_name")
assert class_name is not None, "Model config does not contain a _class_name attribute. Only diffusers format is supported."
config.pop("_diffusers_version")
scheduler_cls, _ = ModelRegistry.resolve_model_cls(class_name)
scheduler = scheduler_cls(**config)
if fastvideo_args.flow_shift is not None:
scheduler.set_shift(fastvideo_args.flow_shift)
if fastvideo_args.timesteps_scale is not None:
scheduler.set_timesteps_scale(fastvideo_args.timesteps_scale)
if fastvideo_args.pipeline_config.flow_shift is not None:
scheduler.set_shift(fastvideo_args.pipeline_config.flow_shift)
if fastvideo_args.pipeline_config.timesteps_scale is not None:
scheduler.set_timesteps_scale(
fastvideo_args.pipeline_config.timesteps_scale)
return scheduler
@@ -529,7 +523,7 @@ class PipelineComponentLoader:
component_model_path: Path to the component model
transformers_or_diffusers: Whether the module is from transformers or diffusers
architecture: Architecture of the component model
fastvideo_args: Inference arguments
pipeline_args: Inference arguments
Returns:
The loaded module
+8 -6
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_replicate", "dp_shards"),
mesh_shape=(hsdp_replicate_dim, hsdp_shard_dim),
mesh_dim_names=("replicate", "shard"),
)
shard_model(model,
cpu_offload=cpu_offload,
+1 -1
View File
@@ -49,7 +49,7 @@ def build_pipeline(fastvideo_args: FastVideoArgs) -> PipelineWithLoRA:
pipeline_architecture)
# instantiate the pipeline
pipeline = pipeline_cls(model_path, fastvideo_args, config)
pipeline = pipeline_cls(model_path, fastvideo_args)
logger.info("Pipeline instantiated")
# pipeline is now initialized and ready to use
@@ -8,13 +8,11 @@ This module defines the base class for pipelines that are composed of multiple s
import argparse
import os
from abc import ABC, abstractmethod
from copy import deepcopy
from typing import Any, Dict, List, Optional, Union, cast
import torch
from fastvideo.v1.configs.pipelines import (PipelineConfig,
get_pipeline_config_cls_for_name)
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.distributed import (
maybe_init_distributed_environment_and_model_parallel)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
@@ -22,7 +20,7 @@ from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import PipelineComponentLoader
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages import PipelineStage
from fastvideo.v1.utils import (maybe_download_model, shallow_asdict,
from fastvideo.v1.utils import (maybe_download_model,
verify_model_config_and_directory)
logger = init_logger(__name__)
@@ -46,24 +44,16 @@ class ComposedPipelineBase(ABC):
# TODO(will): args should support both inference args and training args
def __init__(self,
model_path: str,
fastvideo_args: FastVideoArgs,
config: Optional[Dict[str, Any]] = None,
fastvideo_args: Union[FastVideoArgs, TrainingArgs],
required_config_modules: Optional[List[str]] = None,
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None):
"""
Initialize the pipeline. After __init__, the pipeline should be ready to
use. The pipeline should be stateless and not hold any batch state.
"""
self.fastvideo_args = fastvideo_args
if fastvideo_args.training_mode:
assert isinstance(fastvideo_args, TrainingArgs)
self.training_args = fastvideo_args
assert self.training_args is not None
else:
self.fastvideo_args = fastvideo_args
assert self.fastvideo_args is not None
self.model_path = model_path
self.model_path: str = model_path
self._stages: List[PipelineStage] = []
self._stage_name_mapping: Dict[str, PipelineStage] = {}
@@ -74,13 +64,6 @@ class ComposedPipelineBase(ABC):
raise NotImplementedError(
"Subclass must set _required_config_modules")
if config is None:
# Load configuration
logger.info("Loading pipeline configuration...")
self.config = self._load_config(model_path)
else:
self.config = config
maybe_init_distributed_environment_and_model_parallel(
fastvideo_args.tp_size, fastvideo_args.sp_size)
@@ -89,6 +72,8 @@ class ComposedPipelineBase(ABC):
self.modules = self.load_modules(fastvideo_args, loaded_modules)
if fastvideo_args.training_mode:
assert isinstance(fastvideo_args, TrainingArgs)
self.training_args = fastvideo_args
assert self.training_args is not None
if self.training_args.log_validation:
self.initialize_validation_pipeline(self.training_args)
@@ -127,39 +112,16 @@ class ComposedPipelineBase(ABC):
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
"""
config = None
# 1. If users provide a pipeline config, it will override the default pipeline config
if isinstance(pipeline_config, PipelineConfig):
config = pipeline_config
else:
config_cls = get_pipeline_config_cls_for_name(model_path)
if config_cls is not None:
config = config_cls()
if isinstance(pipeline_config, str):
config.load_from_json(pipeline_config)
# 2. If users also provide some kwargs, it will override the pipeline config.
# The user kwargs shouldn't contain model config parameters!
if config is None:
logger.warning("No config found for model %s, using default config",
model_path)
config_args = kwargs
else:
config_args = shallow_asdict(config)
config_args.update(kwargs)
if args is None or args.inference_mode:
fastvideo_args = FastVideoArgs(model_path=model_path, **config_args)
fastvideo_args.model_path = model_path
for key, value in config_args.items():
setattr(fastvideo_args, key, value)
kwargs['model_path'] = model_path
fastvideo_args = FastVideoArgs.from_kwargs(kwargs)
else:
assert args is not None, "args must be provided for training mode"
fastvideo_args = TrainingArgs.from_cli_args(args)
# TODO(will): fix this so that its not so ugly
fastvideo_args.model_path = model_path
for key, value in config_args.items():
for key, value in kwargs.items():
setattr(fastvideo_args, key, value)
fastvideo_args.use_cpu_offload = False
@@ -170,7 +132,7 @@ class ComposedPipelineBase(ABC):
# use FSDP2's MixedPrecisionPolicy to set the precision for the
# fwd, bwd, and other operations' precision.
# fastvideo_args.precision = fastvideo_args.master_weight_type
assert fastvideo_args.master_weight_type == 'fp32', 'only fp32 is supported for training'
assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
# assert fastvideo_args.precision == 'fp32', 'only fp32 is supported for training'
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
@@ -250,20 +212,21 @@ class ComposedPipelineBase(ABC):
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
"""
logger.info("Loading pipeline modules from config: %s", self.config)
modules_config = deepcopy(self.config)
model_index = self._load_config(self.model_path)
logger.info("Loading pipeline modules from config: %s", model_index)
# remove keys that are not pipeline modules
modules_config.pop("_class_name")
modules_config.pop("_diffusers_version")
model_index.pop("_class_name")
model_index.pop("_diffusers_version")
# some sanity checks
assert len(
modules_config
model_index
) > 1, "model_index.json must contain at least one pipeline module"
for module_name in self.required_config_modules:
if module_name not in modules_config:
if module_name not in model_index:
raise ValueError(
f"model_index.json must contain a {module_name} module")
@@ -273,7 +236,7 @@ class ComposedPipelineBase(ABC):
modules = {}
for module_name, (transformers_or_diffusers,
architecture) in modules_config.items():
architecture) in model_index.items():
if module_name not in required_modules:
logger.info("Skipping module %s", module_name)
continue
+7 -6
View File
@@ -36,16 +36,17 @@ class LoRAPipeline(ComposedPipelineBase):
"transformer"].config.arch_config.exclude_lora_layers
self.convert_to_lora_layers()
if self.fastvideo_args.lora_path is not None:
if self.fastvideo_args.pipeline_config.lora_path is not None:
self.set_lora_adapter(
self.fastvideo_args.lora_nickname, # type: ignore
self.fastvideo_args.lora_path)
self.fastvideo_args.pipeline_config.
lora_nickname, # type: ignore
self.fastvideo_args.pipeline_config.lora_path)
def is_target_layer(self, module_name: str) -> bool:
if self.fastvideo_args.lora_target_names is None:
if self.fastvideo_args.pipeline_config.lora_target_names is None:
return True
return any(target_name in module_name
for target_name in self.fastvideo_args.lora_target_names)
return any(target_name in module_name for target_name in
self.fastvideo_args.pipeline_config.lora_target_names)
def convert_to_lora_layers(self) -> None:
"""
@@ -67,6 +67,7 @@ class ForwardBatch:
# Latent tensors
latents: Optional[torch.Tensor] = None
raw_latent_shape: Optional[torch.Tensor] = None
noise_pred: Optional[torch.Tensor] = None
image_latent: Optional[torch.Tensor] = None
@@ -15,7 +15,7 @@ from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_i2v
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.pipelines.preprocess_pipeline_base import (
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
@@ -87,7 +87,7 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
"clip_feature_dtype": "",
})
return record
return record # type: ignore
EntryClass = PreprocessPipeline_I2V
@@ -6,7 +6,7 @@ This module contains an implementation of the T2V Data Preprocessing pipeline
using the modular pipeline architecture.
"""
from fastvideo.v1.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.v1.pipelines.preprocess_pipeline_base import (
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
@@ -1,37 +1,38 @@
import argparse
import json
import os
import torch
import torch.distributed as dist
from fastvideo.v1.logger import init_logger
from fastvideo.v1.utils import maybe_download_model, shallow_asdict
from fastvideo.v1.distributed import maybe_init_distributed_environment_and_model_parallel, get_world_size
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo import PipelineConfig
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_i2v import PreprocessPipeline_I2V
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_t2v import PreprocessPipeline_T2V
from fastvideo.v1.configs.models.vaes import WanVAEConfig
from fastvideo.v1.distributed import (
get_world_size, maybe_init_distributed_environment_and_model_parallel)
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_i2v import (
PreprocessPipeline_I2V)
from fastvideo.v1.pipelines.preprocess.preprocess_pipeline_t2v import (
PreprocessPipeline_T2V)
from fastvideo.v1.utils import maybe_download_model
logger = init_logger(__name__)
def main(args):
args.model_path = maybe_download_model(args.model_path)
maybe_init_distributed_environment_and_model_parallel(args.tp_size, args.sp_size)
def main(args) -> None:
args.model_path = maybe_download_model(args.model_path)
maybe_init_distributed_environment_and_model_parallel(1, 1)
num_gpus = int(os.environ["WORLD_SIZE"])
assert num_gpus == 1, "Only support 1 GPU"
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
"use_cpu_offload": False,
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
}
pipeline_config_args = shallow_asdict(pipeline_config)
pipeline_config_args.update(kwargs)
fastvideo_args = FastVideoArgs(model_path=args.model_path,
num_gpus=get_world_size(),
**pipeline_config_args,
)
pipeline_config.update_config_from_dict(kwargs)
fastvideo_args = FastVideoArgs(
model_path=args.model_path,
num_gpus=get_world_size(),
pipeline_config=pipeline_config,
)
PreprocessPipeline = PreprocessPipeline_I2V if args.preprocess_task == "i2v" else PreprocessPipeline_T2V
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
@@ -49,7 +50,8 @@ if __name__ == "__main__":
"--dataloader_num_workers",
type=int,
default=1,
help="Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
help=
"Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process.",
)
parser.add_argument(
"--preprocess_video_batch_size",
@@ -63,18 +65,15 @@ if __name__ == "__main__":
default=8,
help="Batch size (per device) for the training dataloader.",
)
parser.add_argument(
"--samples_per_file",
type=int,
default=64
)
parser.add_argument(
"--flush_frequency",
type=int,
default=256,
help="how often to save to parquet files"
)
parser.add_argument("--num_latent_t", type=int, default=28, help="Number of latent timesteps.")
parser.add_argument("--samples_per_file", type=int, default=64)
parser.add_argument("--flush_frequency",
type=int,
default=256,
help="how often to save to parquet files")
parser.add_argument("--num_latent_t",
type=int,
default=28,
help="Number of latent timesteps.")
parser.add_argument("--max_height", type=int, default=480)
parser.add_argument("--max_width", type=int, default=848)
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
@@ -88,22 +87,18 @@ if __name__ == "__main__":
parser.add_argument("--speed_factor", type=float, default=1.0)
parser.add_argument("--drop_short_ratio", type=float, default=1.0)
# text encoder & vae & diffusion model
parser.add_argument("--text_encoder_name", type=str, default="google/t5-v1_1-xxl")
parser.add_argument("--text_encoder_name",
type=str,
default="google/t5-v1_1-xxl")
parser.add_argument("--cache_dir", type=str, default="./cache_dir")
parser.add_argument("--cfg", type=float, default=0.0)
parser.add_argument(
"--output_dir",
type=str,
default=None,
help="The output directory where the model predictions and checkpoints will be written.",
)
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***."),
help=
"The output directory where the model predictions and checkpoints will be written.",
)
args = parser.parse_args()
main(args)
main(args)
+3 -2
View File
@@ -54,7 +54,8 @@ class DecodingStage(PipelineStage):
image = latents
else:
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (vae_dtype != torch.float32
) and not fastvideo_args.disable_autocast
@@ -77,7 +78,7 @@ class DecodingStage(PipelineStage):
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.vae_tiling:
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
+15 -12
View File
@@ -21,7 +21,7 @@ from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.platforms import _Backend
from fastvideo.v1.platforms import AttentionBackendEnum
st_attn_available = False
if importlib.util.find_spec("st_attn") is not None:
@@ -54,10 +54,11 @@ class DenoisingStage(PipelineStage):
self.attn_backend = get_attn_backend(
head_size=attn_head_size,
dtype=torch.float16, # TODO(will): hack
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN,
_Backend.VIDEO_SPARSE_ATTN,
_Backend.FLASH_ATTN,
_Backend.TORCH_SDPA) # hack
supported_attention_backends=(
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA
) # hack
)
def forward(
@@ -194,13 +195,15 @@ class DenoisingStage(PipelineStage):
# Prepare inputs for transformer
t_expand = t.repeat(latent_model_input.shape[0])
guidance_expand = (torch.tensor(
[fastvideo_args.embedded_cfg_scale] *
latent_model_input.shape[0],
dtype=torch.float32,
device=get_torch_device(),
).to(target_dtype) * 1000.0 if fastvideo_args.embedded_cfg_scale
is not None else None)
guidance_expand = (
torch.tensor(
[fastvideo_args.pipeline_config.embedded_cfg_scale] *
latent_model_input.shape[0],
dtype=torch.float32,
device=get_torch_device(),
).to(target_dtype) *
1000.0 if fastvideo_args.pipeline_config.embedded_cfg_scale
is not None else None)
# Predict noise residual
with torch.autocast(device_type="cuda",
+3 -2
View File
@@ -75,7 +75,8 @@ class EncodingStage(PipelineStage):
dtype=torch.float32)
# Setup VAE precision
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.vae_precision]
vae_dtype = PRECISION_TO_TYPE[
fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
@@ -83,7 +84,7 @@ class EncodingStage(PipelineStage):
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.vae_tiling:
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
@@ -75,10 +75,10 @@ class LatentPreparationStage(PipelineStage):
batch_size,
self.transformer.num_channels_latents,
num_frames,
height //
fastvideo_args.vae_config.arch_config.spatial_compression_ratio,
width //
fastvideo_args.vae_config.arch_config.spatial_compression_ratio,
height // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
width // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
)
# Validate generator if it's a list
@@ -103,6 +103,7 @@ class LatentPreparationStage(PipelineStage):
# Update batch with prepared latents
batch.latents = latents
batch.raw_latent_shape = latents.shape
return batch
@@ -119,9 +120,9 @@ class LatentPreparationStage(PipelineStage):
The batch with adjusted video length.
"""
video_length = batch.num_frames
use_temporal_scaling_frames = fastvideo_args.vae_config.use_temporal_scaling_frames
use_temporal_scaling_frames = fastvideo_args.pipeline_config.vae_config.use_temporal_scaling_frames
if use_temporal_scaling_frames:
temporal_scale_factor = fastvideo_args.vae_config.arch_config.temporal_compression_ratio
temporal_scale_factor = fastvideo_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
latent_num_frames = (video_length - 1) // temporal_scale_factor + 1
else: # stepvideo only
latent_num_frames = video_length // 17 * 3
@@ -29,9 +29,9 @@ class StepvideoPromptEncodingStage(PipelineStage):
def forward(self, batch: ForwardBatch, fastvideo_args) -> ForwardBatch:
prompts = [batch.prompt + fastvideo_args.pos_magic]
prompts = [batch.prompt + fastvideo_args.pipeline_config.pos_magic]
bs = len(prompts)
prompts += [fastvideo_args.neg_magic] * bs
prompts += [fastvideo_args.pipeline_config.neg_magic] * bs
with set_forward_context(current_timestep=0, attn_metadata=None):
y, y_mask = self.stepllm(prompts)
clip_emb, _ = self.clip(prompts)
@@ -53,13 +53,13 @@ class TextEncodingStage(PipelineStage):
"""
assert len(self.tokenizers) == len(self.text_encoders)
assert len(self.text_encoders) == len(
fastvideo_args.text_encoder_configs)
fastvideo_args.pipeline_config.text_encoder_configs)
for tokenizer, text_encoder, encoder_config, preprocess_func, postprocess_func in zip(
self.tokenizers, self.text_encoders,
fastvideo_args.text_encoder_configs,
fastvideo_args.preprocess_text_funcs,
fastvideo_args.postprocess_text_funcs):
fastvideo_args.pipeline_config.text_encoder_configs,
fastvideo_args.pipeline_config.preprocess_text_funcs,
fastvideo_args.pipeline_config.postprocess_text_funcs):
if fastvideo_args.use_cpu_offload:
text_encoder = text_encoder.to(get_torch_device())
@@ -9,7 +9,6 @@ using the modular pipeline architecture.
"""
import os
from copy import deepcopy
from typing import Any, Dict
import torch
@@ -102,21 +101,21 @@ class StepVideoPipeline(LoRAPipeline, ComposedPipelineBase):
"""
Load the modules from the config.
"""
logger.info("Loading pipeline modules from config: %s", self.config)
modules_config = deepcopy(self.config)
model_index = self._load_config(self.model_path)
logger.info("Loading pipeline modules from config: %s", model_index)
# remove keys that are not pipeline modules
modules_config.pop("_class_name")
modules_config.pop("_diffusers_version")
model_index.pop("_class_name")
model_index.pop("_diffusers_version")
# some sanity checks
assert len(
modules_config
model_index
) > 1, "model_index.json must contain at least one pipeline module"
required_modules = ["transformer", "scheduler", "vae"]
for module_name in required_modules:
if module_name not in modules_config:
if module_name not in model_index:
raise ValueError(
f"model_index.json must contain a {module_name} module")
logger.info("Diffusers config passed sanity checks")
@@ -124,7 +123,7 @@ class StepVideoPipeline(LoRAPipeline, ComposedPipelineBase):
# all the component models used by the pipeline
modules = {}
for module_name, (transformers_or_diffusers,
architecture) in modules_config.items():
architecture) in model_index.items():
component_model_path = os.path.join(self.model_path, module_name)
module = PipelineComponentLoader.load_module(
module_name=module_name,
@@ -32,7 +32,7 @@ class WanImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
+2 -2
View File
@@ -32,7 +32,7 @@ class WanPipeline(LoRAPipeline, ComposedPipelineBase):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# We use UniPCMScheduler from Wan2.1 official repo, not the one in diffusers.
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
@@ -75,7 +75,7 @@ class WanValidationPipeline(ComposedPipelineBase):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
+1 -1
View File
@@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Optional
from fastvideo.v1.logger import init_logger
# imported by other files, do not remove
from fastvideo.v1.platforms.interface import _Backend # noqa: F401
from fastvideo.v1.platforms.interface import AttentionBackendEnum # noqa: F401
from fastvideo.v1.platforms.interface import Platform, PlatformEnum
from fastvideo.v1.utils import resolve_obj_by_qualname
+32 -15
View File
@@ -13,8 +13,9 @@ from typing_extensions import ParamSpec
import fastvideo.v1.envs as envs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.platforms.interface import (DeviceCapability, Platform,
PlatformEnum, _Backend)
from fastvideo.v1.platforms.interface import (AttentionBackendEnum,
DeviceCapability, Platform,
PlatformEnum)
from fastvideo.v1.utils import import_pynvml
logger = init_logger(__name__)
@@ -106,75 +107,85 @@ class CudaPlatformBase(Platform):
return float(torch.cuda.max_memory_allocated(device))
@classmethod
def get_attn_backend_cls(cls, selected_backend: Optional[_Backend],
def get_attn_backend_cls(cls,
selected_backend: Optional[AttentionBackendEnum],
head_size: int, dtype: torch.dtype) -> str:
# TODO(will): maybe come up with a more general interface for local attention
# if distributed is False, we always try to use Flash attn
logger.info("Trying FASTVIDEO_ATTENTION_BACKEND=%s",
envs.FASTVIDEO_ATTENTION_BACKEND)
if selected_backend == _Backend.SLIDING_TILE_ATTN:
if selected_backend == AttentionBackendEnum.SLIDING_TILE_ATTN:
try:
from st_attn import sliding_tile_attention # noqa: F401
from fastvideo.v1.attention.backends.sliding_tile_attn import ( # noqa: F401
SlidingTileAttentionBackend)
logger.info("Using Sliding Tile Attention backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "SLIDING_TILE_ATTN"
return "fastvideo.v1.attention.backends.sliding_tile_attn.SlidingTileAttentionBackend"
except ImportError as e:
logger.info(e)
logger.info(
"Sliding Tile Attention backend is not installed. Fall back to Flash Attention."
)
elif selected_backend == _Backend.SAGE_ATTN:
elif selected_backend == AttentionBackendEnum.SAGE_ATTN:
try:
from sageattention import sageattn # noqa: F401
from fastvideo.v1.attention.backends.sage_attn import ( # noqa: F401
SageAttentionBackend)
logger.info("Using Sage Attention backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "SAGE_ATTN"
return "fastvideo.v1.attention.backends.sage_attn.SageAttentionBackend"
except ImportError as e:
logger.info(e)
logger.info(
"Sage Attention backend is not installed. Fall back to Flash Attention."
)
elif selected_backend == _Backend.VIDEO_SPARSE_ATTN:
elif selected_backend == AttentionBackendEnum.VIDEO_SPARSE_ATTN:
try:
from vsa import block_sparse_attn # noqa: F401
from fastvideo.v1.attention.backends.video_sparse_attn import ( # noqa: F401
VideoSparseAttentionBackend)
logger.info("Using Video Sparse Attention backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "VIDEO_SPARSE_ATTN"
return "fastvideo.v1.attention.backends.video_sparse_attn.VideoSparseAttentionBackend"
except ImportError as e:
logger.info(e)
logger.info(
"Video Sparse Attention backend is not installed. Fall back to Flash Attention."
)
elif selected_backend == _Backend.TORCH_SDPA:
elif selected_backend == AttentionBackendEnum.TORCH_SDPA:
logger.info("Using Torch SDPA backend.")
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
elif selected_backend == _Backend.FLASH_ATTN or selected_backend is None:
elif selected_backend == AttentionBackendEnum.FLASH_ATTN or selected_backend is None:
pass
elif selected_backend:
raise ValueError(f"Invalid attention backend for {cls.device_name}")
target_backend = _Backend.FLASH_ATTN
target_backend = AttentionBackendEnum.FLASH_ATTN
if not cls.has_device_capability(80):
logger.info(
"Cannot use FlashAttention-2 backend for Volta and Turing "
"GPUs.")
target_backend = _Backend.TORCH_SDPA
target_backend = AttentionBackendEnum.TORCH_SDPA
elif dtype not in (torch.float16, torch.bfloat16):
logger.info(
"Cannot use FlashAttention-2 backend for dtype other than "
"torch.float16 or torch.bfloat16.")
target_backend = _Backend.TORCH_SDPA
target_backend = AttentionBackendEnum.TORCH_SDPA
# FlashAttn is valid for the model, checking if the package is
# installed.
if target_backend == _Backend.FLASH_ATTN:
if target_backend == AttentionBackendEnum.FLASH_ATTN:
try:
import flash_attn # noqa: F401
@@ -187,19 +198,25 @@ class CudaPlatformBase(Platform):
logger.info(
"Cannot use FlashAttention-2 backend for head size %d.",
head_size)
target_backend = _Backend.TORCH_SDPA
target_backend = AttentionBackendEnum.TORCH_SDPA
except ImportError:
logger.info("Cannot use FlashAttention-2 backend because the "
"flash_attn package is not found. "
"Make sure that flash_attn was built and installed "
"(on by default).")
target_backend = _Backend.TORCH_SDPA
target_backend = AttentionBackendEnum.TORCH_SDPA
if target_backend == _Backend.TORCH_SDPA:
if target_backend == AttentionBackendEnum.TORCH_SDPA:
logger.info("Using Torch SDPA backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "TORCH_SDPA"
return "fastvideo.v1.attention.backends.sdpa.SDPABackend"
logger.info("Using Flash Attention backend.")
# Overwrite with the actual backend
envs.FASTVIDEO_ATTENTION_BACKEND = "FLASH_ATTN"
return "fastvideo.v1.attention.backends.flash_attn.FlashAttentionBackend"
@classmethod
+3 -2
View File
@@ -13,7 +13,7 @@ from fastvideo.v1.logger import init_logger
logger = init_logger(__name__)
class _Backend(enum.Enum):
class AttentionBackendEnum(enum.Enum):
FLASH_ATTN = enum.auto()
SLIDING_TILE_ATTN = enum.auto()
TORCH_SDPA = enum.auto()
@@ -88,7 +88,8 @@ class Platform:
return self._enum == PlatformEnum.CUDA
@classmethod
def get_attn_backend_cls(cls, selected_backend: Optional[_Backend],
def get_attn_backend_cls(cls,
selected_backend: Optional[AttentionBackendEnum],
head_size: int, dtype: torch.dtype) -> str:
"""Get the attention backend class of a device."""
return ""
-2
View File
@@ -1,8 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import pytest
import torch.distributed as dist
import pytest
import torch
import numpy as np
@@ -10,6 +10,7 @@ from transformers import AutoConfig
from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
load_tokenizer)
# from fastvideo.v1.models.hunyuan.text_encoder import load_text_encoder, load_tokenizer
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
@@ -40,8 +41,7 @@ def test_clip_encoder():
- Produce nearly identical outputs for the same input prompts
"""
args = FastVideoArgs(model_path="openai/clip-vit-large-patch14",
text_encoder_precisions=("fp16",),
text_encoder_configs=(CLIPTextConfig(),))
pipeline_config=PipelineConfig(text_encoder_configs=(CLIPTextConfig(),), text_encoder_precisions=("fp16",)))
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
logger.info("Loading models from %s", args.model_path)
@@ -8,6 +8,7 @@ from transformers import AutoConfig
from fastvideo.models.hunyuan.text_encoder import (load_text_encoder,
load_tokenizer)
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
@@ -40,8 +41,7 @@ def test_llama_encoder():
- Produce nearly identical outputs for the same input prompts
"""
args = FastVideoArgs(model_path="meta-llama/Llama-2-7b-hf",
text_encoder_precisions=("fp16",),
text_encoder_configs=(LlamaConfig(),))
pipeline_config=PipelineConfig(text_encoder_configs=(LlamaConfig(),), text_encoder_precisions=("fp16",)))
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
@@ -6,6 +6,7 @@ import pytest
import torch
from transformers import AutoConfig, AutoTokenizer, UMT5EncoderModel
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import TextEncoderLoader
@@ -39,7 +40,8 @@ def test_t5_encoder():
precision).to(device).eval()
tokenizer = AutoTokenizer.from_pretrained(TOKENIZER_PATH)
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH, text_encoder_configs=(T5Config(),), text_encoder_precisions=(precision_str,))
args = FastVideoArgs(model_path=TEXT_ENCODER_PATH, pipeline_config=PipelineConfig(text_encoder_configs=(T5Config(),), text_encoder_precisions=(precision_str,)))
loader = TextEncoderLoader()
model2 = loader.load(TEXT_ENCODER_PATH, "", args)
@@ -0,0 +1,177 @@
import os
from pathlib import Path
from huggingface_hub import snapshot_download
import shutil
import subprocess
import sys
from fastvideo.v1.tests.ssim.test_inference_similarity import compute_video_ssim_torchvision
# Import the training pipeline
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
NUM_NODES = "1"
MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
# preprocessing
DATA_DIR = "data"
LOCAL_RAW_DATA_DIR = Path(os.path.join(DATA_DIR, "cats"))
NUM_GPUS_PER_NODE_PREPROCESSING = "1"
PREPROCESSING_ENTRY_FILE_PATH = "fastvideo/v1/pipelines/preprocess/v1_preprocess.py"
LOCAL_PREPROCESSED_DATA_DIR = Path(os.path.join(DATA_DIR, "cats_preprocessed_data"))
# training
NUM_GPUS_PER_NODE_TRAINING = "4"
TRAINING_ENTRY_FILE_PATH = "fastvideo/v1/training/wan_training_pipeline.py"
LOCAL_TRAINING_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "combined_parquet_dataset")
LOCAL_VALIDATION_DATA_DIR = os.path.join(LOCAL_PREPROCESSED_DATA_DIR, "validation_parquet_dataset")
LOCAL_OUTPUT_DIR = Path(os.path.join(DATA_DIR, "outputs"))
def download_data():
# create the data dir if it doesn't exist
data_dir = Path(DATA_DIR)
if data_dir.exists():
print(f"Removing existing data directory at {data_dir}")
shutil.rmtree(data_dir)
print(f"Creating data directory at {data_dir}")
os.makedirs(data_dir)
print(f"Downloading raw dataset to {LOCAL_RAW_DATA_DIR}...")
try:
result = snapshot_download(
repo_id="wlsaidhi/cats-overfit-merged",
local_dir=str(LOCAL_RAW_DATA_DIR),
repo_type="dataset",
resume_download=True,
token=os.environ.get("HF_TOKEN"), # In case authentication is needed
)
print(f"Download completed successfully. Files downloaded to: {result}")
# Verify the download
if not LOCAL_RAW_DATA_DIR.exists():
raise RuntimeError(f"Download appeared to succeed but {LOCAL_RAW_DATA_DIR} does not exist")
# List downloaded files
print("Downloaded files:")
for file in LOCAL_RAW_DATA_DIR.rglob("*"):
if file.is_file():
print(f" - {file.relative_to(LOCAL_RAW_DATA_DIR)}")
except Exception as e:
print(f"Error during download: {str(e)}")
raise
def run_preprocessing():
# Run torchrun command
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE_PREPROCESSING,
PREPROCESSING_ENTRY_FILE_PATH,
"--model_path", MODEL_PATH,
"--data_merge_path", os.path.join(LOCAL_RAW_DATA_DIR, "merge_1_sample.txt"),
"--preprocess_video_batch_size", "1",
"--max_height", "480",
"--max_width", "832",
"--num_frames", "77",
"--dataloader_num_workers", "0",
"--output_dir", LOCAL_PREPROCESSED_DATA_DIR,
"--train_fps", "16",
"--validation_prompt_txt", os.path.join(LOCAL_RAW_DATA_DIR, "validation_prompt_1_sample.txt"),
"--samples_per_file", "1",
"--flush_frequency", "1",
"--video_length_tolerance_range", "5",
"--dataset", "t2v",
]
process = subprocess.run(cmd, check=True)
def run_training():
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE_TRAINING,
TRAINING_ENTRY_FILE_PATH,
"--model_path", MODEL_PATH,
"--inference_mode", "False",
"--pretrained_model_name_or_path", MODEL_PATH,
"--data_path", LOCAL_TRAINING_DATA_DIR,
"--validation_prompt_dir", LOCAL_VALIDATION_DATA_DIR,
"--train_batch_size", "1",
"--num_latent_t", "8",
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
"--sp_size", "4",
"--tp_size", "4",
"--hsdp_replicate_dim", "1",
"--hsdp_shard_dim", "4",
"--num_gpus", NUM_GPUS_PER_NODE_TRAINING,
"--train_sp_batch_size", "1",
"--dataloader_num_workers", "10",
"--gradient_accumulation_steps", "1",
"--max_train_steps", "901",
"--learning_rate", "1e-5",
"--mixed_precision", "bf16",
"--checkpointing_steps", "6000",
"--validation_steps", "100",
"--validation_sampling_steps", "50",
"--log_validation",
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--cfg", "0.0",
"--output_dir", LOCAL_OUTPUT_DIR,
"--tracker_project_name", "wan_finetune_overfit_ci",
"--num_height", "480",
"--num_width", "832",
"--num_frames", "81",
"--validation_guidance_scale", "1.0",
"--num_euler_timesteps", "50",
"--multi_phased_distill_schedule", "4000-1",
"--weight_decay", "0.01",
"--not_apply_cfg_solver",
"--dit_precision", "fp32",
"--max_grad_norm", "1.0",
]
print(f"Running training with command: {cmd}")
process = subprocess.run(cmd, check=True)
def test_e2e_overfit_single_sample():
os.environ["WANDB_MODE"] = "online"
download_data()
run_preprocessing()
run_training()
reference_video_file = os.path.join(os.path.dirname(__file__), "reference_video_1_sample_v0.mp4")
print(f"reference_video_file: {reference_video_file}")
final_validation_video_file = os.path.join(LOCAL_OUTPUT_DIR, "validation_step_900_inference_steps_50_video_0.mp4")
print(f"final_validation_video_file: {final_validation_video_file}")
# Ensure both files exist
assert os.path.exists(reference_video_file), f"Reference video not found at {reference_video_file}"
assert os.path.exists(final_validation_video_file), f"Validation video not found at {final_validation_video_file}"
# Compute SSIM
mean_ssim, min_ssim, max_ssim = compute_video_ssim_torchvision(
reference_video_file,
final_validation_video_file,
use_ms_ssim=True # Using MS-SSIM for better quality assessment
)
print("\n===== SSIM Results for Step 900 Validation =====")
print(f"Mean MS-SSIM: {mean_ssim:.4f}")
print(f"Min MS-SSIM: {min_ssim:.4f}")
print(f"Max MS-SSIM: {max_ssim:.4f}")
assert max_ssim > 0.5, f"Max SSIM is below 0.5: {max_ssim}"
if __name__ == "__main__":
test_e2e_overfit_single_sample()
@@ -15,7 +15,7 @@ def setup_args():
parser = argparse.ArgumentParser(description='T5 Encoder Test')
parser.add_argument('--model_path', type=str, default="google/umt5-xxl")
parser.add_argument(
'--precision',
'--dit-precision',
type=str,
default="float32",
help='Precision to use for the model (float32, float16, bfloat16)')
@@ -0,0 +1 @@
{"_timestamp":1.7496170016478686e+09,"validation_videos_8_steps":{"_type":"videos","count":5,"videos":[{"sha256":"d81aa715df0c3ba4db8b242bc10442844453f3684a45dc55b89ecd27b2414fc7","size":429642,"path":"media/videos/validation_videos_8_steps_0_d81aa715df0c3ba4db8b.mp4","_type":"video-file"},{"sha256":"cd72a3d513eca6b41b03b80e6fa044ce7219c35e969d2ca20b9cb48c91e585c6","size":477837,"path":"media/videos/validation_videos_8_steps_0_cd72a3d513eca6b41b03.mp4","_type":"video-file"},{"_type":"video-file","sha256":"43d47c211a69bf0be3544738e76e7d8fa58bb108c00cb21dc34eaad3c6ce7cc3","size":409419,"path":"media/videos/validation_videos_8_steps_0_43d47c211a69bf0be354.mp4"},{"_type":"video-file","sha256":"ea674ec9e200bc97563c9d87d9dc07110c3f42c0ab7277dd237ae96dd8f90a10","size":333966,"path":"media/videos/validation_videos_8_steps_0_ea674ec9e200bc97563c.mp4"},{"sha256":"d81aa715df0c3ba4db8b242bc10442844453f3684a45dc55b89ecd27b2414fc7","size":429642,"path":"media/videos/validation_videos_8_steps_0_d81aa715df0c3ba4db8b.mp4","_type":"video-file"}],"captions":false},"step_time":2.5065076276659966,"_wandb":{"runtime":53},"learning_rate":1e-06,"_step":5,"_runtime":53.172758961,"grad_norm":5.65625,"train_loss":0.3915919363498688,"avg_step_time":2.8116052336990833}
@@ -0,0 +1,140 @@
import os
import sys
import subprocess
from pathlib import Path
import torch.distributed.elastic.multiprocessing.errors as errors
from torch.distributed.elastic.multiprocessing.errors import record
from torch.utils.data import DataLoader
import torch
import pytest
import wandb
import json
from huggingface_hub import snapshot_download
# Import the training pipeline
sys.path.append(str(Path(__file__).parent.parent.parent.parent.parent))
from fastvideo.v1.training.wan_training_pipeline import main
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.v1.utils import FlexibleArgumentParser
wandb_name = "test_training_loss"
reference_wandb_summary_file = "fastvideo/v1/tests/training/reference_wandb_summary.json"
NUM_NODES = "1"
NUM_GPUS_PER_NODE = "4"
def run_worker():
"""Worker function that will be run on each GPU"""
# Create and populate args
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
# Set the arguments as they are in finetune_v1_test.sh
args = parser.parse_args([
"--model_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--inference_mode", "False",
"--pretrained_model_name_or_path", "Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
"--cache_dir", "/home/.cache",
"--data_path", "data/crush-smol_parq/combined_parquet_dataset",
"--validation_prompt_dir", "data/crush-smol_parq/validation_parquet_dataset",
"--train_batch_size", "2",
"--num_latent_t", "4",
"--num_gpus", "4",
"--sp_size", "4",
"--tp_size", "4",
"--hsdp_replicate_dim", "1",
"--hsdp_shard_dim", "4",
"--train_sp_batch_size", "1",
"--dataloader_num_workers", "1",
"--gradient_accumulation_steps", "2",
"--max_train_steps", "5",
"--learning_rate", "1e-6",
"--mixed_precision", "bf16",
"--checkpointing_steps", "30",
"--validation_steps", "10",
"--validation_sampling_steps", "8",
"--log_validation",
"--checkpoints_total_limit", "3",
"--allow_tf32",
"--ema_start_step", "0",
"--cfg", "0.0",
"--output_dir", "data/wan_finetune_test",
"--tracker_project_name", "wan_finetune_ci",
"--wandb_run_name", wandb_name,
"--num_height", "480",
"--num_width", "832",
"--num_frames", "81",
"--flow_shift", "3",
"--validation_guidance_scale", "1.0",
"--num_euler_timesteps", "50",
"--multi_phased_distill_schedule", "4000-1",
"--weight_decay", "0.01",
"--not_apply_cfg_solver",
"--dit_precision", "fp32",
"--max_grad_norm", "1.0"
])
# Call the main training function
main(args)
def test_distributed_training():
"""Test the distributed training setup"""
os.environ["WANDB_MODE"] = "offline"
data_dir = Path("data/crush-smol_parq")
if not data_dir.exists():
print(f"Downloading test dataset to {data_dir}...")
snapshot_download(
repo_id="PY007/crush-smol",
local_dir=str(data_dir),
repo_type="dataset",
local_dir_use_symlinks=False
)
# Get the current file path
current_file = Path(__file__).resolve()
# Run torchrun command
cmd = [
"torchrun",
"--nnodes", NUM_NODES,
"--nproc_per_node", NUM_GPUS_PER_NODE,
str(current_file)
]
process = subprocess.run(cmd, check=True)
summary_file = "fastvideo/v1/tests/training/reference_wandb_summary.json"
reference_wandb_summary = json.load(open(reference_wandb_summary_file))
wandb_summary = json.load(open(summary_file))
fields_and_thresholds = {
'avg_step_time': 1.0,
'grad_norm': 0.1,
'step_time': 0.5,
'train_loss': 0.001
}
failures = []
for field, threshold in fields_and_thresholds.items():
ref_value = reference_wandb_summary[field]
current_value = wandb_summary[field]
diff = abs(ref_value - current_value)
print(f"INFO: {field}, diff: {diff}, threshold: {threshold}, reference: {ref_value}, current: {current_value}")
if diff > threshold:
failures.append(f"FAILED: {field} difference {diff} exceeds threshold of {threshold} (reference: {ref_value}, current: {current_value})")
if failures:
raise AssertionError("\n".join(failures))
if __name__ == "__main__":
if os.environ.get("LOCAL_RANK") is not None:
# We're being run by torchrun
run_worker()
else:
# We're being run directly
test_distributed_training()
@@ -6,6 +6,7 @@ import os
import pytest
import torch
from fastvideo.v1.configs.pipelines.base import PipelineConfig
from fastvideo.v1.distributed.parallel_state import (
get_sp_parallel_rank,
get_sp_world_size)
@@ -62,10 +63,8 @@ def test_hunyuanvideo_distributed():
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
use_cpu_offload=False,
precision=precision_str)
pipeline_config=PipelineConfig(dit_config=HunyuanVideoConfig(), dit_precision=precision_str))
args.device = torch.device(f"cuda:{LOCAL_RANK}")
args.dit_config = HunyuanVideoConfig()
args.check_fastvideo_args()
loader = TransformerLoader()
model = loader.load(TRANSFORMER_PATH, "", args)
@@ -6,6 +6,7 @@ import pytest
import torch
from diffusers import WanTransformer3DModel
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.forward_context import set_forward_context
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
@@ -34,10 +35,8 @@ def test_wan_transformer():
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
use_cpu_offload=False,
precision=precision_str)
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
args.dit_config = WanVideoConfig()
args.check_fastvideo_args()
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, "", args).to(device, dtype=precision)
+2 -2
View File
@@ -7,6 +7,7 @@ import pytest
import torch
from safetensors.torch import load_file
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.logger import init_logger
# from fastvideo.v1.models.vaes.hunyuanvae import (
# AutoencoderKLHunyuanVideo as MyHunyuanVAE)
@@ -36,9 +37,8 @@ def test_hunyuan_vae():
device = torch.device("cuda:0")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=VAE_PATH, vae_precision=precision_str)
args = FastVideoArgs(model_path=VAE_PATH, pipeline_config=PipelineConfig(vae_config=HunyuanVAEConfig(), vae_precision=precision_str))
args.device = device
args.vae_config = HunyuanVAEConfig()
loader = VAELoader()
model = loader.load(VAE_PATH, "", args)
+2 -2
View File
@@ -6,6 +6,7 @@ import pytest
import torch
from diffusers import AutoencoderKLWan
from fastvideo.v1.configs.pipelines import PipelineConfig
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.loader.component_loader import VAELoader
@@ -29,9 +30,8 @@ def test_wan_vae():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=VAE_PATH, vae_precision=precision_str)
args = FastVideoArgs(model_path=VAE_PATH, pipeline_config=PipelineConfig(vae_config=WanVAEConfig(), vae_precision=precision_str))
args.device = device
args.vae_config = WanVAEConfig()
loader = VAELoader()
model2 = loader.load(VAE_PATH, "", args)
+117 -129
View File
@@ -4,7 +4,7 @@ import math
import os
import traceback
from abc import ABC, abstractmethod
from typing import Any, Dict, Iterator
from typing import Any, Dict, Iterator, List
import imageio
import numpy as np
@@ -16,7 +16,7 @@ from einops import rearrange
from torchdata.stateful_dataloader import StatefulDataLoader
from fastvideo.v1.configs.sample import SamplingParam
from fastvideo.v1.dataset.parquet_datasets import ParquetVideoTextDataset
from fastvideo.v1.dataset import build_parquet_map_style_dataloader
from fastvideo.v1.distributed import (get_sp_group, get_torch_device,
get_world_group)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
@@ -93,22 +93,15 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
last_epoch=self.init_steps - 1,
)
self.train_dataset = ParquetVideoTextDataset(
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
training_args.data_path,
batch_size=training_args.train_batch_size,
cfg_rate=training_args.cfg,
num_latent_t=training_args.num_latent_t)
self.train_dataloader = StatefulDataLoader(
self.train_dataset,
batch_size=training_args.train_batch_size,
num_workers=training_args.
dataloader_num_workers, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
pin_memory_device=f"cuda:{torch.cuda.current_device()}",
drop_last=True)
training_args.train_batch_size,
num_data_workers=training_args.dataloader_num_workers,
drop_last=True,
text_padding_length=training_args.pipeline_config.
text_encoder_configs[0].arch_config.
text_len, # type: ignore[attr-defined]
seed=training_args.seed)
self.noise_scheduler = noise_scheduler
@@ -128,7 +121,9 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
if self.global_rank == 0:
project = training_args.tracker_project_name or "fastvideo"
wandb.init(project=project, config=training_args)
wandb.init(project=project,
config=training_args,
name=training_args.wandb_run_name)
@abstractmethod
def initialize_validation_pipeline(self, training_args: TrainingArgs):
@@ -172,138 +167,131 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
# Prepare validation prompts
logger.info('fastvideo_args.validation_prompt_dir: %s',
training_args.validation_prompt_dir)
validation_dataset = ParquetVideoTextDataset(
validation_dataset, validation_dataloader = build_parquet_map_style_dataloader(
training_args.validation_prompt_dir,
batch_size=1,
cfg_rate=training_args.cfg,
num_latent_t=training_args.num_latent_t,
validation=True)
num_data_workers=0,
drop_last=False,
drop_first_row=sampling_param.negative_prompt is not None,
cfg_rate=training_args.cfg)
if sampling_param.negative_prompt:
_, negative_prompt_embeds, negative_prompt_attention_mask, _ = validation_dataset.get_validation_negative_prompt(
)
validation_dataloader = StatefulDataLoader(
validation_dataset,
batch_size=1,
num_workers=5, # Reduce number of workers to avoid memory issues
prefetch_factor=2,
shuffle=False,
pin_memory=True,
pin_memory_device=f"cuda:{torch.cuda.current_device()}",
drop_last=False)
transformer.eval()
# Process each validation prompt
videos = []
captions = []
for _, embeddings, masks, infos in validation_dataloader:
caption = infos['caption']
captions.extend(caption)
prompt_embeds = embeddings.to(get_torch_device())
prompt_attention_mask = masks.to(get_torch_device())
validation_steps = training_args.validation_sampling_steps.split(",")
validation_steps = [int(step) for step in validation_steps]
validation_steps = [step for step in validation_steps if step > 0]
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8,
sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# Process each validation prompt for each validation step
for num_inference_steps in validation_steps:
step_videos: List[np.ndarray] = []
step_captions: List[str | None] = []
temporal_compression_factor = training_args.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
for _, embeddings, masks, infos in validation_dataloader:
step_captions.extend([None]) # TODO(peiyuan): add caption
prompt_embeds = embeddings.to(get_torch_device())
prompt_attention_mask = masks.to(get_torch_device())
# Prepare batch for validation
batch = ForwardBatch(
data_type="video",
latents=None,
seed=validation_seed, # Use deterministic seed
generator=torch.Generator(
device="cpu").manual_seed(validation_seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
# make sure we use the same height, width, and num_frames as the training pipeline
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
# TODO(will): validation_sampling_steps and
# validation_guidance_scale are actually passed in as a list of
# values, like "10,20,30". The validation should be run for each
# combination of values.
# num_inference_steps=fastvideo_args.validation_sampling_steps,
num_inference_steps=sampling_param.num_inference_steps,
# guidance_scale=fastvideo_args.validation_guidance_scale,
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
eta=0.0,
)
# Calculate sizes
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8,
sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
# Run validation inference
with torch.no_grad(), torch.autocast("cuda", dtype=torch.bfloat16):
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
# Re-enable gradients for training
transformer.requires_grad_(True)
transformer.train()
# Prepare batch for validation
batch = ForwardBatch(
data_type="video",
latents=None,
seed=validation_seed, # Use deterministic seed
generator=torch.Generator(
device="cpu").manual_seed(validation_seed),
prompt_embeds=[prompt_embeds],
prompt_attention_mask=[prompt_attention_mask],
negative_prompt_embeds=[negative_prompt_embeds],
negative_attention_mask=[negative_prompt_attention_mask],
height=training_args.num_height,
width=training_args.num_width,
num_frames=num_frames,
num_inference_steps=
num_inference_steps, # Use the current validation step
guidance_scale=sampling_param.guidance_scale,
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
if self.rank_in_sp_group != 0:
continue
# Run validation inference
with torch.no_grad(), torch.autocast("cuda",
dtype=torch.bfloat16):
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
videos.append(frames)
if self.rank_in_sp_group != 0:
continue
# Log validation results
world_group = get_world_group()
num_sp_groups = world_group.world_size // self.sp_group.world_size
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
step_videos.append(frames)
# Only sp_group leaders (rank_in_sp_group == 0) need to send their
# results to global rank 0
if self.rank_in_sp_group == 0:
if self.global_rank == 0:
# Global rank 0 collects results from all sp_group leaders
all_videos = videos # Start with own results
all_captions = captions
# Log validation results for this step
world_group = get_world_group()
num_sp_groups = world_group.world_size // self.sp_group.world_size
# Receive from other sp_group leaders
for sp_group_idx in range(1, num_sp_groups):
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
recv_videos = world_group.recv_object(src=src_rank)
recv_captions = world_group.recv_object(src=src_rank)
all_videos.extend(recv_videos)
all_captions.extend(recv_captions)
# Only sp_group leaders (rank_in_sp_group == 0) need to send their
# results to global rank 0
if self.rank_in_sp_group == 0:
if self.global_rank == 0:
# Global rank 0 collects results from all sp_group leaders
all_videos = step_videos # Start with own results
all_captions = step_captions
video_filenames = []
for i, (video,
caption) in enumerate(zip(all_videos, all_captions)):
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
f"validation_step_{global_step}_video_{i}.mp4")
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
# Receive from other sp_group leaders
for sp_group_idx in range(1, num_sp_groups):
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
recv_videos = world_group.recv_object(src=src_rank)
recv_captions = world_group.recv_object(src=src_rank)
all_videos.extend(recv_videos)
all_captions.extend(recv_captions)
logs = {
"validation_videos": [
wandb.Video(filename, caption=caption) for filename,
caption in zip(video_filenames, all_captions)
]
}
wandb.log(logs, step=global_step)
else:
# Other sp_group leaders send their results to global rank 0
world_group.send_object(videos, dst=0)
world_group.send_object(captions, dst=0)
video_filenames = []
for i, (video,
caption) in enumerate(zip(all_videos,
all_captions)):
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
f"validation_step_{global_step}_inference_steps_{num_inference_steps}_video_{i}.mp4"
)
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
logs = {
f"validation_videos_{num_inference_steps}_steps": [
wandb.Video(filename, caption=caption)
for filename, caption in zip(
video_filenames, all_captions)
]
}
wandb.log(logs, step=global_step)
else:
# Other sp_group leaders send their results to global rank 0
world_group.send_object(step_videos, dst=0)
world_group.send_object(step_captions, dst=0)
# Re-enable gradients for training
transformer.train()
gc.collect()
torch.cuda.empty_cache()
+16
View File
@@ -9,8 +9,11 @@ import torch
import torch.distributed as dist
import torch.distributed.checkpoint as dcp
import torch.distributed.checkpoint.stateful
from einops import rearrange
from safetensors.torch import save_file
from fastvideo.v1.distributed.parallel_state import (get_sp_parallel_rank,
get_sp_world_size)
from fastvideo.v1.logger import init_logger
from fastvideo.v1.training.checkpointing_utils import (ModelWrapper,
OptimizerWrapper,
@@ -252,6 +255,19 @@ def normalize_dit_input(model_type, latents, args=None) -> torch.Tensor:
raise NotImplementedError(f"model_type {model_type} not supported")
def shard_latents_across_sp(latents: torch.Tensor,
num_latent_t: int) -> torch.Tensor:
sp_world_size = get_sp_world_size()
rank_in_sp_group = get_sp_parallel_rank()
latents = latents[:, :, :num_latent_t]
if sp_world_size > 1:
latents = rearrange(latents,
"b c (n s) h w -> b c n s h w",
n=sp_world_size).contiguous()
latents = latents[:, :, rank_in_sp_group, :, :, :]
return latents
def clip_grad_norm_while_handling_failing_dtensor_cases(
parameters: Union[torch.Tensor, List[torch.Tensor]],
max_norm: float,
+58 -16
View File
@@ -1,4 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
import importlib.util
import random
import sys
import time
@@ -10,6 +11,9 @@ import torch
from diffusers import FlowMatchEulerDiscreteScheduler
from tqdm.auto import tqdm
import fastvideo.v1.envs as envs
from fastvideo.v1.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadata)
from fastvideo.v1.distributed import (cleanup_dist_env_and_memory, get_sp_group,
get_torch_device, get_world_group)
from fastvideo.v1.fastvideo_args import FastVideoArgs, TrainingArgs
@@ -23,10 +27,14 @@ from fastvideo.v1.training.training_pipeline import TrainingPipeline
from fastvideo.v1.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, get_sigmas, load_checkpoint,
normalize_dit_input, save_checkpoint)
normalize_dit_input, save_checkpoint, shard_latents_across_sp)
import wandb # isort: skip
vsa_available = False
if importlib.util.find_spec("vsa") is not None:
vsa_available = True
logger = init_logger(__name__)
# Manual gradient checking flag - set to True to enable gradient verification
@@ -41,7 +49,7 @@ class WanTrainingPipeline(TrainingPipeline):
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.flow_shift)
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
@@ -54,16 +62,19 @@ class WanTrainingPipeline(TrainingPipeline):
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
args_copy.vae_config.load_encoder = False
args_copy.pipeline_config.vae_config.load_encoder = False
validation_pipeline = WanValidationPipeline.from_pretrained(
training_args.model_path,
args=None,
inference_mode=True,
loaded_modules={"transformer": self.get_module("transformer")})
loaded_modules={"transformer": self.get_module("transformer")},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus)
self.validation_pipeline = validation_pipeline
def train_one_step(
def train_one_step( # type: ignore[override]
self,
transformer,
model_type,
@@ -80,6 +91,8 @@ class WanTrainingPipeline(TrainingPipeline):
logit_mean,
logit_std,
mode_scale,
patch_size,
current_vsa_sparsity,
) -> tuple[float, float]:
assert self.training_args is not None
self.modules["transformer"].requires_grad_(True)
@@ -99,16 +112,20 @@ class WanTrainingPipeline(TrainingPipeline):
# Get first batch of new epoch
batch = next(self.train_loader_iter)
latents, encoder_hidden_states, encoder_attention_mask, infos = batch
latents, encoder_hidden_states, encoder_attention_mask, _ = batch
# logger.info("rank: %s, caption: %s",
# self.rank,
# infos['caption'],
# local_main_process_only=False)
# TODO(will): don't hardcode bfloat16
latents = latents.to(get_torch_device(), dtype=torch.bfloat16)
encoder_hidden_states = encoder_hidden_states.to(
get_torch_device(), dtype=torch.bfloat16)
latents = shard_latents_across_sp(
latents, num_latent_t=self.training_args.num_latent_t)
dit_seq_shape = [
latents.shape[2] // patch_size[0],
latents.shape[3] // patch_size[1],
latents.shape[4] // patch_size[2]
]
latents = normalize_dit_input(model_type, latents)
batch_size = latents.shape[0]
noise = torch.randn_like(latents)
@@ -148,8 +165,17 @@ class WanTrainingPipeline(TrainingPipeline):
[1000.0],
device=noisy_model_input.device,
dtype=torch.bfloat16)
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
attn_metadata = VideoSparseAttentionMetadata(
current_timestep=timesteps,
dit_seq_shape=dit_seq_shape,
VSA_sparsity=current_vsa_sparsity)
else:
attn_metadata = None
with set_forward_context(current_timestep=timesteps,
attn_metadata=None):
attn_metadata=attn_metadata):
model_pred = transformer(**input_kwargs)
if precondition_outputs:
@@ -273,11 +299,23 @@ class WanTrainingPipeline(TrainingPipeline):
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info("GPU memory usage before train_one_step: %s MB",
gpu_memory_usage)
logger.info("VSA validation sparsity: %s",
self.training_args.VSA_sparsity)
self._log_validation(self.transformer, self.training_args, 1)
if vsa_available:
vsa_sparsity = self.training_args.VSA_sparsity
vsa_decay_rate = self.training_args.VSA_decay_rate
vsa_decay_interval_steps = self.training_args.VSA_decay_interval_steps
for step in range(self.init_steps + 1,
self.training_args.max_train_steps + 1):
start_time = time.perf_counter()
if vsa_available:
current_decay_times = min(step // vsa_decay_interval_steps,
vsa_sparsity // vsa_decay_rate)
current_vsa_sparsity = current_decay_times * vsa_decay_rate
else:
current_vsa_sparsity = 0.0
loss, grad_norm = self.train_one_step(
self.transformer,
# args.model_type,
@@ -295,10 +333,9 @@ class WanTrainingPipeline(TrainingPipeline):
self.training_args.logit_mean,
self.training_args.logit_std,
self.training_args.mode_scale,
self.training_args.pipeline_config.dit_config.patch_size,
current_vsa_sparsity,
)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info("GPU memory usage after train_one_step: %s MB",
gpu_memory_usage)
step_time = time.perf_counter() - start_time
step_times.append(step_time)
@@ -325,6 +362,7 @@ class WanTrainingPipeline(TrainingPipeline):
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
"vsa_sparsity": current_vsa_sparsity,
},
step=step,
)
@@ -337,7 +375,11 @@ class WanTrainingPipeline(TrainingPipeline):
self.sp_group.barrier()
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
self._log_validation(self.transformer, self.training_args, step)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info("GPU memory usage after validation: %s MB",
gpu_memory_usage)
wandb.finish()
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir,
self.training_args.max_train_steps, self.optimizer,
+2 -2
View File
@@ -238,7 +238,7 @@ class FlexibleArgumentParser(argparse.ArgumentParser):
"serve,chat,complete",
"facebook/opt-12B",
'--port', '12323',
'--tensor-parallel-size', '4',
'--tp-size', '4',
'-tp', '2'
]
```
@@ -291,7 +291,7 @@ class FlexibleArgumentParser(argparse.ArgumentParser):
returns:
processed_args: list[str] = [
'--port': '12323',
'--tensor-parallel-size': '4',
'--tp-size': '4',
'--vae-config.load-encoder': 'false',
'--vae-config.load-decoder': 'true'
]
+29
View File
@@ -0,0 +1,29 @@
from abc import ABC, abstractmethod
from typing import Dict
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.pipelines import ComposedPipelineBase, build_pipeline
class WorkflowBase(ABC):
pipeline_configs: Dict[str, FastVideoArgs] = {}
pipelines: Dict[str, ComposedPipelineBase] = {}
def __init__(self, fastvideo_args: FastVideoArgs):
self.fastvideo_args = fastvideo_args
def register_pipelines(self, pipeline_configs: Dict[str, FastVideoArgs]):
self.pipeline_configs.update(pipeline_configs)
def load_pipelines(self):
for pipeline_name, pipeline_config in self.pipeline_configs.items():
pipeline = build_pipeline(pipeline_config)
self.pipelines[pipeline_name] = pipeline
@abstractmethod
def get_components(self):
pass
@abstractmethod
def run(self):
pass
+112 -113
View File
@@ -1,138 +1,137 @@
# SPDX-License-Identifier: Apache-2.0
import argparse
import json
import os
import time
from multiprocessing import Pool, cpu_count
from pathlib import Path
import cv2
import torchvision
from tqdm import tqdm
def get_video_info(video_path, prompt_text):
"""Extract video information using OpenCV and corresponding prompt text"""
cap = cv2.VideoCapture(str(video_path))
def get_video_info(video_path):
"""Get video information using torchvision."""
# Read video tensor (T, C, H, W)
video_tensor, _, info = torchvision.io.read_video(str(video_path),
output_format="TCHW",
pts_unit="sec")
if not cap.isOpened():
print(f"Error: Could not open video {video_path}")
return None
num_frames = video_tensor.shape[0]
height = video_tensor.shape[2]
width = video_tensor.shape[3]
fps = info.get("video_fps", 0)
duration = num_frames / fps if fps > 0 else 0
# Get video properties
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
fps = cap.get(cv2.CAP_PROP_FPS)
frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
duration = frame_count / fps if fps > 0 else 0
cap.release()
# Extract name
_, _, videos_dir, video_name = str(video_path).split("/")
return {
"path": video_path.name,
"path": str(video_name),
"resolution": {
"width": width,
"height": height
},
"size": os.path.getsize(video_path),
"fps": fps,
"duration": duration,
"cap": [prompt_text]
"num_frames": num_frames
}
def read_prompt_file(prompt_path):
"""Read and return the content of a prompt file"""
try:
with open(prompt_path, 'r', encoding='utf-8') as f:
return f.read().strip()
except Exception as e:
print(f"Error reading prompt file {prompt_path}: {e}")
return None
def prepare_dataset_json(folder_path,
output_name="videos2caption.json",
num_workers=None) -> None:
"""Prepare dataset information from a folder containing videos and prompt.txt."""
folder_path = Path(folder_path)
# Read prompt file
prompt_file = folder_path / "prompt.txt"
if not prompt_file.exists():
raise FileNotFoundError(f"prompt.txt not found in {folder_path}")
with open(prompt_file) as f:
prompts = [line.strip() for line in f.readlines() if line.strip()]
# Read videos file
videos_file = folder_path / "videos.txt"
if not videos_file.exists():
raise FileNotFoundError(f"videos.txt not found in {folder_path}")
with open(videos_file) as f:
video_paths = [line.strip() for line in f.readlines() if line.strip()]
if len(prompts) != len(video_paths):
raise ValueError(
f"Number of prompts ({len(prompts)}) does not match number of videos ({len(video_paths)})"
)
# Prepare arguments for multiprocessing
process_args = [folder_path / video_path for video_path in video_paths]
# Determine number of workers
if num_workers is None:
num_workers = max(1, cpu_count() - 1) # Leave one CPU free
# Process videos in parallel
start_time = time.time()
with Pool(num_workers) as pool:
results = list(
tqdm(pool.imap(get_video_info, process_args),
total=len(process_args),
desc="Processing videos",
unit="video"))
# Combine results with prompts
dataset_info = []
for result, prompt in zip(results, prompts):
result["cap"] = [prompt]
dataset_info.append(result)
# Calculate total processing time
total_time = time.time() - start_time
total_videos = len(dataset_info)
avg_time_per_video = total_time / total_videos if total_videos > 0 else 0
print("\nProcessing completed:")
print(f"Total videos processed: {total_videos}")
print(f"Total time: {total_time:.2f} seconds")
print(f"Average time per video: {avg_time_per_video:.2f} seconds")
# Save to JSON file
output_file = folder_path / output_name
with open(output_file, 'w') as f:
json.dump(dataset_info, f, indent=2)
# Create merge.txt
merge_file = folder_path / "merge.txt"
with open(merge_file, 'w') as f:
f.write(f"{folder_path}/videos,{output_file}\n")
print(f"Dataset information saved to {output_file}")
print(f"Merge file created at {merge_file}")
def process_videos_and_prompts(video_dir_path, prompt_dir_path, verbose=False):
"""Process videos and their corresponding prompt files
Args:
video_dir_path (str): Path to directory containing video files
prompt_dir_path (str): Path to directory containing prompt files
verbose (bool): Whether to print verbose processing information
"""
video_dir = Path(video_dir_path)
prompt_dir = Path(prompt_dir_path)
processed_data = []
# Ensure directories exist
if not video_dir.exists() or not prompt_dir.exists():
print(f"Error: One or both directories do not exist:\nVideos: {video_dir}\nPrompts: {prompt_dir}")
return []
# Process each video file
for video_file in video_dir.glob('*.mp4'):
video_name = video_file.stem
prompt_file = prompt_dir / f"{video_name}.txt"
# Check if corresponding prompt file exists
if not prompt_file.exists():
print(f"Warning: No prompt file found for video {video_name}")
continue
# Read prompt content
prompt_text = read_prompt_file(prompt_file)
if prompt_text is None:
continue
# Process video and add to results
video_info = get_video_info(video_file, prompt_text)
if video_info:
processed_data.append(video_info)
return processed_data
def save_results(processed_data, output_path):
"""Save processed data to JSON file
Args:
processed_data (list): List of processed video information
output_path (str): Full path for output JSON file
"""
output_path = Path(output_path)
# Create parent directories if they don't exist
output_path.parent.mkdir(parents=True, exist_ok=True)
with open(output_path, 'w', encoding='utf-8') as f:
json.dump(processed_data, f, indent=2, ensure_ascii=False)
return output_path
def parse_args():
"""Parse command line arguments"""
import argparse
parser = argparse.ArgumentParser(description='Process videos and their corresponding prompt files')
parser.add_argument('--video_dir', '-v', required=True, help='Directory containing video files')
parser.add_argument('--prompt_dir', '-p', required=True, help='Directory containing prompt text files')
parser.add_argument('--output_path',
'-o',
required=True,
help='Full path for output JSON file (e.g., /path/to/output/videos2caption.json)')
parser.add_argument('--verbose', action='store_true', help='Print verbose processing information')
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description='Prepare video dataset information in JSON format')
parser.add_argument(
'--data_folder',
type=str,
required=True,
help='Path to the folder containing videos and prompt.txt')
parser.add_argument(
'--output',
type=str,
default='videos2caption.json',
help='Name of the output JSON file (default: videos2caption.json)')
parser.add_argument('--workers',
type=int,
default=32,
help='Number of worker processes (default: 16)')
return parser.parse_args()
if __name__ == "__main__":
# Parse command line arguments
args = parse_args()
# Process videos and prompts
processed_videos = process_videos_and_prompts(args.video_dir, args.prompt_dir, args.verbose)
if processed_videos:
# Save results
output_path = save_results(processed_videos, args.output_path)
print(f"\nProcessed {len(processed_videos)} videos")
print(f"Results saved to: {output_path}")
# Print example of processed data
print("\nExample of processed video info:")
print(json.dumps(processed_videos[0], indent=2))
else:
print("No videos were processed successfully")
prepare_dataset_json(args.data_folder, args.output, args.workers)
+18 -21
View File
@@ -2,35 +2,35 @@ export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
DATA_DIR=data/HD-Mixkit-Finetune-Wan/combined_parquet_dataset
VALIDATION_DIR=data/HD-Mixkit-Finetune-Wan/validation_parquet_dataset
NUM_GPUS=1
DATA_DIR=[your data dir]
VALIDATION_DIR=[your validation dir]
NUM_GPUS=4
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
CHECKPOINT_PATH="$DATA_DIR/outputs/wan_finetune/checkpoint-5"
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
fastvideo/v1/training/wan_training_pipeline.py\
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache"\
--data_path "$DATA_DIR"\
--validation_prompt_dir "$VALIDATION_DIR"\
--train_batch_size=1\
--num_latent_t 4 \
--sp_size $NUM_GPUS \
--tp_size $NUM_GPUS \
--train_batch_size=4 \
--num_latent_t 20 \
--sp_size 4 \
--tp_size 4 \
--hsdp_replicate_dim 1 \
--hsdp_shard_dim 4 \
--num_gpus $NUM_GPUS \
--train_sp_batch_size 1\
--dataloader_num_workers 5\
--gradient_accumulation_steps=1\
--max_train_steps=120 \
--learning_rate=1e-6\
--dataloader_num_workers 10\
--gradient_accumulation_steps=1 \
--max_train_steps=5000 \
--learning_rate=1e-5\
--mixed_precision="bf16"\
--checkpointing_steps=50 \
--validation_steps 20\
--checkpointing_steps=6000 \
--validation_steps 50\
--validation_sampling_steps "2,4,8" \
--log_validation \
--checkpoints_total_limit 3\
@@ -42,13 +42,10 @@ torchrun --nnodes 1 --nproc_per_node $NUM_GPUS\
--num_height 480 \
--num_width 832 \
--num_frames 81 \
--shift 3 \
--validation_guidance_scale "1.0" \
--num_euler_timesteps 50 \
--multi_phased_distill_schedule "4000-1" \
--weight_decay 0.01 \
--not_apply_cfg_solver \
--master_weight_type "fp32" \
--max_grad_norm 1.0 \
# --resume_from_checkpoint "$CHECKPOINT_PATH"
--dit_precision "fp32" \
--max_grad_norm 1.0
+62
View File
@@ -0,0 +1,62 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export WANDB_API_KEY='your_wandb_api_key'
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
export TRITON_CACHE_DIR=/tmp/triton_cache
DATA_DIR=~/train/
VALIDATION_DIR=~latents/test/
NUM_GPUS=8
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
CHECKPOINT_PATH="$DATA_DIR/outputs/wan_finetune/checkpoint-5"
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun --nnodes 1 --nproc_per_node $NUM_GPUS \
fastvideo/v1/training/wan_training_pipeline.py \
--model_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--cache_dir "/home/ray/.cache" \
--data_path "$DATA_DIR" \
--validation_prompt_dir "$VALIDATION_DIR" \
--train_batch_size 1 \
--num_latent_t 16 \
--sp_size 1 \
--tp_size 1 \
--num_gpus $NUM_GPUS \
--hsdp_replicate_dim $NUM_GPUS \
--hsdp-shard-dim 1 \
--train_sp_batch_size 1 \
--dataloader_num_workers 4 \
--gradient_accumulation_steps 8 \
--max_train_steps 30000 \
--learning_rate 1e-5 \
--mixed_precision "bf16" \
--checkpointing_steps 6000 \
--validation_steps 100 \
--validation_sampling_steps "50" \
--log_validation \
--checkpoints_total_limit 3 \
--allow_tf32 \
--ema_start_step 0 \
--cfg 0.0 \
--output_dir "$DATA_DIR/outputs/wan_finetune" \
--tracker_project_name VSA_finetune \
--num_height 448 \
--num_width 832 \
--num_frames 61 \
--flow_shift 3 \
--validation_guidance_scale "5.0" \
--num_euler_timesteps 50 \
--master_weight_type "fp32" \
--dit_precision "fp32" \
--weight_decay 0.01 \
--max_grad_norm 1.0 \
--VSA_decay_sparsity 0.9 \
--VSA_decay_rate 0.03 \
--VSA_decay_interval_steps 30 \
--VSA_val_sparsity 0.9
# --resume_from_checkpoint "$CHECKPOINT_PATH"
@@ -2,12 +2,12 @@
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-I2V-14B-480P-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="finetrainers/crush-smol/merge.txt"
OUTPUT_DIR="crush-smol_preprocess"
VALIDATION_PATH="assets/prompt.txt"
DATA_MERGE_PATH="mini_i2v_dataset/crush-smol_raw/merge.txt"
OUTPUT_DIR="mini_i2v_dataset/crush-smol_preprocessed"
VALIDATION_PATH="mini_i2v_dataset/crush-smol_raw/validation.txt"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/data_preprocess/preprocess.py \
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 8 \
@@ -19,6 +19,6 @@ torchrun --nproc_per_node=$GPU_NUM \
--model_type $MODEL_TYPE \
--train_fps 16 \
--validation_prompt_txt $VALIDATION_PATH \
--samples_per_file 16 \
--flush_frequency 32 \
--samples_per_file 8 \
--flush_frequency 8 \
--preprocess_task "i2v"
@@ -7,7 +7,7 @@ OUTPUT_DIR="data/crush-smol/latents"
VALIDATION_PATH="assets/prompt.txt"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/data_preprocess/preprocess.py \
fastvideo/v1/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 1 \