Compare commits

..
Author SHA1 Message Date
RandNMR73 da5ca94091 preprocessing text 2025-09-06 05:56:13 +00:00
66 changed files with 1468 additions and 2996 deletions
-23
View File
@@ -176,26 +176,3 @@ steps:
- TEST_TYPE=precision_vsa
agents:
queue: "default"
- path:
- "csrc/attn/vmoba_attn/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Precision Tests VMoBA"
env:
- TEST_TYPE=precision_vmoba
agents:
queue: "default"
- path:
- "csrc/attn/vmoba_attn/vmoba/**"
- "fastvideo/attention/backends/vmoba.py"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Inference Tests VMoBA"
env:
- TEST_TYPE=inference_vmoba
agents:
queue: "default"
-9
View File
@@ -109,15 +109,6 @@ case "$TEST_TYPE" in
log "Running distillation DMD tests..."
MODAL_COMMAND="$MODAL_ENV WANDB_API_KEY=$WANDB_API_KEY python3 -m modal run $MODAL_TEST_FILE::run_distill_dmd_tests"
;;
# run_inference_tests_vmoba
"inference_vmoba")
log "Running V-MoBA inference tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_inference_tests_vmoba"
;;
"precision_vmoba")
log "Running V-MoBA precision tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_precision_tests_vmoba"
;;
*)
log "Error: Unknown test type: $TEST_TYPE"
exit 1
-32
View File
@@ -1,32 +0,0 @@
# Attention Kernel Used in FastVideo
## VMoBA: Mixture-of-Block Attention for Video Diffusion Models (VMoBA)
### Installation
Please ensure that you have installed FlashAttention version **2.7.1 or higher**, as some interfaces have changed in recent releases.
### Usage
You can use `moba_attn_varlen` in the following ways:
**Install from source:**
```bash
python setup.py install
```
**Import after installation:**
```python
from vmoba import moba_attn_varlen
```
**Or import directly from the project root:**
```python
from csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
```
### Verify if you have successfully installed
```bash
python csrc/attn/vmoba_attn/vmoba/vmoba.py
```
-24
View File
@@ -1,24 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from setuptools import find_packages, setup
PACKAGE_NAME = "vmoba"
VERSION = "0.0.0"
AUTHOR = "JianzongWu"
DESCRIPTION = "VMoBA: Mixture-of-Block Attention for Video Diffusion Models"
URL = "https://github.com/KwaiVGI/VMoBA"
setup(
name=PACKAGE_NAME,
version=VERSION,
author=AUTHOR,
description=DESCRIPTION,
url=URL,
packages=find_packages(),
classifiers=[
"Programming Language :: Python :: 3",
"License :: OSI Approved :: Apache Software License",
],
python_requires='>=3.12',
install_requires=[]
)
@@ -1,97 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import torch
import pytest
import random
from csrc.attn.vmoba_attn.vmoba import moba_attn_varlen
def generate_test_data(batch_size, total_seqlen, num_heads, head_dim, dtype, device="cuda"):
"""
Generates random data for testing the variable-length attention function.
"""
torch.manual_seed(42)
random.seed(42)
torch.cuda.manual_seed_all(42)
# Generate sequence lengths for each item in the batch
if batch_size > 1:
# Ensure sequence lengths are reasonably distributed
avg_seqlen = total_seqlen // batch_size
seqlens = [random.randint(avg_seqlen // 2, avg_seqlen + avg_seqlen // 2) for _ in range(batch_size - 1)]
remaining_len = total_seqlen - sum(seqlens)
if remaining_len > 0:
seqlens.append(remaining_len)
else: # Adjust if sum exceeds total_seqlen
seqlens.append(avg_seqlen)
current_sum = sum(seqlens)
seqlens[-1] -= (current_sum - total_seqlen)
# Ensure all lengths are positive
seqlens = [max(1, s) for s in seqlens]
# Final adjustment to match total_seqlen
seqlens[-1] += total_seqlen - sum(seqlens)
else:
seqlens = [total_seqlen]
cu_seqlens = torch.tensor([0] + list(torch.cumsum(torch.tensor(seqlens), 0)), device=device, dtype=torch.int32)
max_seqlen = max(seqlens) if seqlens else 0
q = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
k = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
v = torch.randn((total_seqlen, num_heads, head_dim), dtype=dtype, device=device, requires_grad=False)
return q, k, v, cu_seqlens, max_seqlen
@pytest.mark.parametrize("batch_size", [1, 2])
@pytest.mark.parametrize("total_seqlen", [512, 1024])
@pytest.mark.parametrize("num_heads", [8])
@pytest.mark.parametrize("head_dim", [64])
@pytest.mark.parametrize("moba_chunk_size", [64])
@pytest.mark.parametrize("moba_topk", [2, 4])
@pytest.mark.parametrize("select_mode", ["topk", "threshold"])
@pytest.mark.parametrize("threshold_type", ["query_head", "head_global", "overall"])
@pytest.mark.parametrize("dtype", [torch.float32, torch.float16, torch.bfloat16])
def test_moba_attn_varlen_forward(
batch_size, total_seqlen, num_heads, head_dim, moba_chunk_size, moba_topk, select_mode, threshold_type, dtype
):
"""
Tests the forward pass of moba_attn_varlen for basic correctness.
It checks output shape, dtype, and for the presence of NaNs/Infs.
"""
if dtype == torch.float32:
pytest.skip("float32 is not supported in flash attention")
q, k, v, cu_seqlens, max_seqlen = generate_test_data(
batch_size, total_seqlen, num_heads, head_dim, dtype
)
# Ensure chunk size is not larger than the smallest sequence length
min_seqlen = (cu_seqlens[1:] - cu_seqlens[:-1]).min().item()
if moba_chunk_size > min_seqlen:
pytest.skip("moba_chunk_size is larger than the minimum sequence length in the batch")
try:
output = moba_attn_varlen(
q=q,
k=k,
v=v,
cu_seqlens=cu_seqlens,
max_seqlen=max_seqlen,
moba_chunk_size=moba_chunk_size,
moba_topk=moba_topk,
select_mode=select_mode,
threshold_type=threshold_type,
simsum_threshold=0.5, # A reasonable default for threshold mode
)
except Exception as e:
pytest.fail(f"moba_attn_varlen forward pass failed with exception: {e}")
# 1. Check output shape
assert output.shape == q.shape, f"Expected output shape {q.shape}, but got {output.shape}"
# 2. Check output dtype
assert output.dtype == q.dtype, f"Expected output dtype {q.dtype}, but got {output.dtype}"
# 3. Check for NaNs or Infs in the output
assert torch.all(torch.isfinite(output)), "Output contains NaN or Inf values"
-2
View File
@@ -1,2 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from .vmoba import moba_attn_varlen, process_moba_input, process_moba_output
-860
View File
@@ -1,860 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
# Adapt from https://github.com/KwaiVGI/VMoBA/blob/main/src/vmoba.py
import random
import time
import os
import torch
from typing import Tuple
from flash_attn import flash_attn_varlen_func # Use the new flash attention function
from flash_attn.flash_attn_interface import _flash_attn_varlen_forward, _flash_attn_varlen_backward
from functools import lru_cache
from einops import rearrange
@lru_cache(maxsize=16)
def calc_chunks(cu_seqlen, moba_chunk_size):
"""
Calculate chunk boundaries.
For vision tasks we include all chunks (even the last one which might be shorter)
so that every chunk can be selected.
"""
batch_sizes = cu_seqlen[1:] - cu_seqlen[:-1]
batch_num_chunk = (batch_sizes + (moba_chunk_size - 1)) // moba_chunk_size
cu_num_chunk = torch.ones(
batch_num_chunk.numel() + 1,
device=cu_seqlen.device,
dtype=batch_num_chunk.dtype,
)
cu_num_chunk[1:] = batch_num_chunk.cumsum(dim=0)
num_chunk = cu_num_chunk[-1]
chunk_sizes = torch.full(
(num_chunk + 1,), moba_chunk_size, dtype=torch.int32, device=cu_seqlen.device
)
chunk_sizes[0] = 0
batch_last_chunk_size = batch_sizes - (batch_num_chunk - 1) * moba_chunk_size
chunk_sizes[cu_num_chunk[1:]] = batch_last_chunk_size
cu_chunk = chunk_sizes.cumsum(dim=-1, dtype=torch.int32)
chunk_to_batch = torch.zeros(
(num_chunk,), dtype=torch.int32, device=cu_seqlen.device
)
chunk_to_batch[cu_num_chunk[1:-1]] = 1
chunk_to_batch = chunk_to_batch.cumsum(dim=0, dtype=torch.int32)
# Do not filter out any chunk
filtered_chunk_indices = torch.arange(
num_chunk, device=cu_seqlen.device, dtype=torch.int32
)
num_filtered_chunk = num_chunk
return cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch
# --- Threshold Selection Helper Functions ---
def _select_threshold_query_head(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects chunks for each <query, head> pair based on threshold.
Normalization and sorting happen along the chunk dimension (dim=0).
"""
C, H, S = gate.shape
eps = 1e-6
# LSE‐style normalization per <head, query> (across chunks)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
row_min = gate_min_val.amin(dim=0) # (H, S)
row_max = gate_masked.amax(dim=0) # (H, S)
denom = row_max - row_min
denom = torch.where(denom <= eps, torch.ones_like(denom), denom) # avoid divide‑by‑zero
gate_norm = (gate - row_min.unsqueeze(0)) / denom.unsqueeze(0)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) pull out the self‐chunk’s normalized weight for each <head,seq>
self_norm = (gate_norm * gate_self_chunk_mask).sum(dim=0) # (H, S)
# 2) compute how much more normalized weight we need beyond self
total_norm_sum = gate_norm.sum(dim=0) # (H, S)
remain_ratio = simsum_threshold - self_norm / (total_norm_sum + eps) # (H, S)
remain_ratio = torch.clamp(remain_ratio, min=0.0) # if already ≥ thresh, no extra needed
# 3) zero out the self‐chunk in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0
# 4) sort the other chunks by descending norm, per <head,seq>
sorted_norm, sorted_idx = torch.sort(others_norm, descending=True, dim=0) # (C, H, S)
# 5) cumulative‑sum the sorted norms per <head,seq>
cumsum_others = sorted_norm.cumsum(dim=0) # (C, H, S)
# 6) for each <head,seq>, find the smallest k where cumsum_ratio ≥ remain_ratio
ratio = cumsum_others / (total_norm_sum.unsqueeze(0) + eps) # (C, H, S)
cond = ratio >= remain_ratio.unsqueeze(0) # (C, H, S) boolean mask
any_cond = cond.any(dim=0) # (H, S)
# Find the index of the first True value along dim 0. If none, use C-1.
cutoff = torch.where(any_cond, cond.float().argmax(dim=0), torch.full_like(any_cond, fill_value=C - 1)) # (H, S)
# 7) build a mask in sorted order up to that cutoff
idx_range = torch.arange(C, device=gate.device).view(-1, 1, 1) # (C, 1, 1)
sorted_mask = idx_range <= cutoff.unsqueeze(0) # (C, H, S)
# 8) scatter it back to original chunk order
others_mask = torch.zeros_like(gate, dtype=torch.bool)
others_mask.scatter_(0, sorted_idx, sorted_mask)
# 9) finally, include every self‐chunk plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_block(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <query, head> pairs for each block based on threshold.
Normalization and sorting happen across the head and sequence dimensions (dim=1, 2).
"""
C, H, S = gate.shape
HS = H * S
eps = 1e-6
# LSE‐style normalization per block (across heads and queries)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
block_max = gate_masked.amax(dim=(1, 2), keepdim=True) # (C, 1, 1)
block_min = gate_min_val.amin(dim=(1, 2), keepdim=True) # (C, 1, 1)
block_denom = block_max - block_min
block_denom = torch.where(block_denom <= eps, torch.ones_like(block_denom), block_denom) # (C, 1, 1)
gate_norm = (gate - block_min) / block_denom # (C, H, S)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) identify normalized weights of entries that *are* self-chunks (from query perspective)
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
# Sum these weights *per block*
self_norm_sum_per_block = self_norm_entries.sum(dim=(1, 2)) # (C,)
# 2) compute how much more normalized weight each block needs beyond its self-chunk contributions
total_norm_sum_per_block = gate_norm.sum(dim=(1, 2)) # (C,)
remain_ratio = simsum_threshold - self_norm_sum_per_block / (total_norm_sum_per_block + eps) # (C,)
remain_ratio = torch.clamp(remain_ratio, min=0.0) # (C,)
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
# 4) sort the other <head, seq> pairs by descending norm, per block
others_flat = others_norm.contiguous().view(C, HS) # (C, H*S)
sorted_others_flat, sorted_indices_flat = torch.sort(others_flat, dim=1, descending=True) # (C, H*S)
# 5) cumulative‑sum the sorted norms per block
cumsum_others_flat = sorted_others_flat.cumsum(dim=1) # (C, H*S)
# 6) for each block, find the smallest k where cumsum_ratio ≥ remain_ratio
ratio_flat = cumsum_others_flat / (total_norm_sum_per_block.unsqueeze(1) + eps) # (C, H*S)
cond_flat = ratio_flat >= remain_ratio.unsqueeze(1) # (C, H*S) boolean mask
any_cond = cond_flat.any(dim=1) # (C,)
# Find the index of the first True value along dim 1. If none, use HS-1.
cutoff_flat = torch.where(any_cond, cond_flat.float().argmax(dim=1), torch.full_like(any_cond, fill_value=HS - 1)) # (C,)
# 7) build a mask in sorted order up to that cutoff per block
idx_range_flat = torch.arange(HS, device=gate.device).unsqueeze(0) # (1, H*S)
sorted_mask_flat = idx_range_flat <= cutoff_flat.unsqueeze(1) # (C, H*S)
# 8) scatter it back to original <head, seq> order per block
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C, H*S)
others_mask_flat.scatter_(1, sorted_indices_flat, sorted_mask_flat)
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
# 9) finally, include every self‐chunk entry plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_overall(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <chunk, query, head> triplets globally based on threshold.
Normalization and sorting happen across all valid entries.
"""
C, H, S = gate.shape
CHS = C * H * S
eps = 1e-6
# LSE‐style normalization globally across all valid entries
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf) # Use -inf for max
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf) # Use +inf for min
overall_max = gate_masked.max() # scalar
overall_min = gate_min_val.min() # scalar
overall_denom = overall_max - overall_min
overall_denom = torch.where(overall_denom <= eps, torch.tensor(1.0, device=gate.device, dtype=gate.dtype), overall_denom)
gate_norm = (gate - overall_min) / overall_denom # (C, H, S)
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 1) identify normalized weights of entries that *are* self-chunks
self_norm_entries = gate_norm * gate_self_chunk_mask # (C, H, S)
# Sum these weights globally
self_norm_sum_overall = self_norm_entries.sum() # scalar
# 2) compute how much more normalized weight is needed globally beyond self-chunk contributions
total_norm_sum_overall = gate_norm.sum() # scalar
remain_ratio = simsum_threshold - self_norm_sum_overall / (total_norm_sum_overall + eps) # scalar
remain_ratio = torch.clamp(remain_ratio, min=0.0) # scalar
# 3) zero out the self‐chunk entries in a copy, so we only sort “others”
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # Zero out self entries
# 4) sort all other entries by descending norm, globally
others_flat = others_norm.flatten() # (C*H*S,)
valid_others_mask_flat = valid_gate_mask.flatten() & ~gate_self_chunk_mask.flatten() # Mask for valid, non-self entries
# Only sort the valid 'other' entries
valid_others_indices = torch.where(valid_others_mask_flat)[0]
valid_others_values = others_flat[valid_others_indices]
sorted_others_values, sort_perm = torch.sort(valid_others_values, descending=True) # (N_valid_others,)
sorted_original_indices = valid_others_indices[sort_perm] # Original indices in C*H*S space, sorted by value
# 5) cumulative‑sum the sorted valid 'other' norms globally
cumsum_others_values = sorted_others_values.cumsum(dim=0) # (N_valid_others,)
# 6) find the smallest k where cumsum_ratio ≥ remain_ratio globally
ratio_values = cumsum_others_values / (total_norm_sum_overall + eps) # (N_valid_others,)
cond_values = ratio_values >= remain_ratio # (N_valid_others,) boolean mask
any_cond = cond_values.any() # scalar
# Find the index of the first True value in the *sorted* list. If none, use all valid others.
cutoff_idx_in_sorted = torch.where(
any_cond,
cond_values.float().argmax(dim=0),
torch.tensor(len(sorted_others_values) - 1, device=gate.device, dtype=torch.long)
)
# 7) build a mask selecting the top-k others based on the cutoff
# Select the original indices corresponding to the top entries in the sorted list
selected_other_indices = sorted_original_indices[:cutoff_idx_in_sorted + 1]
# 8) create the mask in the original flat shape
others_mask_flat = torch.zeros_like(others_flat, dtype=torch.bool) # (C*H*S,)
if selected_other_indices.numel() > 0: # Check if any 'other' indices were selected
others_mask_flat[selected_other_indices] = True
others_mask = others_mask_flat.view(C, H, S) # (C, H, S)
# 9) finally, include every self‐chunk entry plus all selected others
final_gate_mask = valid_gate_mask & (others_mask | gate_self_chunk_mask)
return final_gate_mask
def _select_threshold_head_global(
gate: torch.Tensor,
valid_gate_mask: torch.Tensor,
gate_self_chunk_mask: torch.Tensor,
simsum_threshold: float
) -> torch.Tensor:
"""
Selects <chunk, query> globally for each head based on threshold.
"""
C, H, S = gate.shape
eps = 1e-6
# 1) LSE‐style normalization per head (across chunks and sequence dims)
gate_masked = torch.where(valid_gate_mask, gate, -torch.inf)
gate_min_val = torch.where(valid_gate_mask, gate, torch.inf)
max_per_head = gate_masked.amax(dim=(0, 2), keepdim=True) # (1, H, 1)
min_per_head = gate_min_val.amin(dim=(0, 2), keepdim=True) # (1, H, 1)
denom = max_per_head - min_per_head
denom = torch.where(denom <= eps, torch.ones_like(denom), denom)
gate_norm = (gate - min_per_head) / denom
gate_norm = torch.where(valid_gate_mask, gate_norm, 0.0) # (C, H, S)
# 2) sum normalized self‐chunk contributions per head
self_norm_sum = (gate_norm * gate_self_chunk_mask).sum(dim=(0, 2)) # (H,)
# 3) total normalized sum per head
total_norm_sum = gate_norm.sum(dim=(0, 2)) # (H,)
# 4) how much more normalized weight needed per head
remain_ratio = simsum_threshold - self_norm_sum / (total_norm_sum + eps) # (H,)
remain_ratio = torch.clamp(remain_ratio, min=0.0)
# 5) zero out self‐chunk entries to focus on "others"
others_norm = gate_norm.clone()
others_norm[gate_self_chunk_mask] = 0.0 # (C, H, S)
# 6) flatten chunk and sequence dims, per head
CS = C * S
others_flat = others_norm.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
valid_flat = (valid_gate_mask & ~gate_self_chunk_mask) \
.permute(1, 0, 2).reshape(H, CS) # (H, C*S)
# 7) vectorized selection of “others” per head
masked_flat = torch.where(valid_flat, others_flat, torch.zeros_like(others_flat))
sorted_vals, sorted_idx = torch.sort(masked_flat, dim=1, descending=True) # (H, C*S)
cumsum_vals = sorted_vals.cumsum(dim=1) # (H, C*S)
ratio_vals = cumsum_vals / (total_norm_sum.unsqueeze(1) + eps) # (H, C*S)
cond = ratio_vals >= remain_ratio.unsqueeze(1) # (H, C*S)
has_cutoff = cond.any(dim=1) # (H,)
default = torch.full((H,), CS - 1, device=gate.device, dtype=torch.long)
cutoff = torch.where(has_cutoff, cond.float().argmax(dim=1), default) # (H,)
idx_range = torch.arange(CS, device=gate.device).unsqueeze(0) # (1, C*S)
sorted_mask = idx_range <= cutoff.unsqueeze(1) # (H, C*S)
selected_flat = torch.zeros_like(valid_flat) # (H, C*S)
selected_flat.scatter_(1, sorted_idx, sorted_mask) # (H, C*S)
# 8) reshape selection mask back to (C, H, S)
others_mask = selected_flat.reshape(H, C, S).permute(1, 0, 2) # (C, H, S)
# 9) include self‐chunks plus selected others, and obey valid mask
final_gate_mask = valid_gate_mask & (gate_self_chunk_mask | others_mask)
return final_gate_mask
class MixedAttention(torch.autograd.Function):
@staticmethod
def forward(
ctx,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
max_seqlen,
moba_chunk_size,
moba_q_sh_indices,
):
ctx.max_seqlen = max_seqlen
ctx.moba_chunk_size = moba_chunk_size
ctx.softmax_scale = softmax_scale = q.shape[-1] ** (-0.5)
# Non-causal self-attention branch
# return out, softmax_lse, S_dmask, rng_state
self_attn_out_sh, self_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
q=q,
k=k,
v=v,
cu_seqlens_q=self_attn_cu_seqlen,
cu_seqlens_k=self_attn_cu_seqlen,
max_seqlen_q=max_seqlen,
max_seqlen_k=max_seqlen,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
)
# MOBA attention branch (non-causal)
moba_attn_out, moba_attn_lse_hs, _, _ = _flash_attn_varlen_forward(
q=moba_q,
k=moba_kv[:, 0],
v=moba_kv[:, 1],
cu_seqlens_q=moba_cu_seqlen_q,
cu_seqlens_k=moba_cu_seqlen_kv,
max_seqlen_q=max_seqlen,
max_seqlen_k=moba_chunk_size,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
)
self_attn_lse_sh = self_attn_lse_hs.t().contiguous()
moba_attn_lse = moba_attn_lse_hs.t().contiguous()
output = torch.zeros((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
output_2d = output.view(-1, q.shape[2])
max_lse_1d = self_attn_lse_sh.view(-1)
max_lse_1d = max_lse_1d.index_reduce(
0, moba_q_sh_indices, moba_attn_lse.view(-1), "amax"
)
self_attn_lse_sh = self_attn_lse_sh - max_lse_1d.view_as(self_attn_lse_sh)
moba_attn_lse = (
moba_attn_lse.view(-1)
.sub(max_lse_1d.index_select(0, moba_q_sh_indices))
.reshape_as(moba_attn_lse)
)
mixed_attn_se_sh = self_attn_lse_sh.exp()
moba_attn_se = moba_attn_lse.exp()
mixed_attn_se_sh.view(-1).index_add_(
0, moba_q_sh_indices, moba_attn_se.view(-1)
)
mixed_attn_lse_sh = mixed_attn_se_sh.log()
# Combine self-attention output
factor = (self_attn_lse_sh - mixed_attn_lse_sh).exp() # [S, H]
self_attn_out_sh = self_attn_out_sh * factor.unsqueeze(-1)
output_2d += self_attn_out_sh.reshape_as(output_2d)
# Combine MOBA attention output
mixed_attn_lse = (
mixed_attn_lse_sh.view(-1)
.index_select(0, moba_q_sh_indices)
.view_as(moba_attn_lse)
)
factor = (moba_attn_lse - mixed_attn_lse).exp() # [S, H]
moba_attn_out = moba_attn_out * factor.unsqueeze(-1)
raw_attn_out = moba_attn_out.view(-1, moba_attn_out.shape[-1])
output_2d.index_add_(0, moba_q_sh_indices, raw_attn_out)
output = output.to(q.dtype)
mixed_attn_lse_sh = mixed_attn_lse_sh + max_lse_1d.view_as(mixed_attn_se_sh)
ctx.save_for_backward(
output,
mixed_attn_lse_sh,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
moba_q_sh_indices,
)
return output
@staticmethod
def backward(ctx, d_output):
max_seqlen = ctx.max_seqlen
moba_chunk_size = ctx.moba_chunk_size
softmax_scale = ctx.softmax_scale
(
output,
mixed_attn_vlse_sh,
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
moba_q_sh_indices,
) = ctx.saved_tensors
d_output = d_output.contiguous()
dq = torch.empty_like(q)
dk = torch.empty_like(k)
dv = torch.empty_like(v)
_ = _flash_attn_varlen_backward(
dout=d_output,
q=q,
k=k,
v=v,
out=output,
softmax_lse=mixed_attn_vlse_sh.t().contiguous(),
dq=dq,
dk=dk,
dv=dv,
cu_seqlens_q=self_attn_cu_seqlen,
cu_seqlens_k=self_attn_cu_seqlen,
max_seqlen_q=max_seqlen,
max_seqlen_k=max_seqlen,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
softcap=0.0,
alibi_slopes=None,
deterministic=True,
window_size_left=-1,
window_size_right=-1
)
headdim = q.shape[-1]
d_moba_output = (
d_output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
)
moba_output = (
output.view(-1, headdim).index_select(0, moba_q_sh_indices).unsqueeze(1)
)
mixed_attn_vlse = (
mixed_attn_vlse_sh.view(-1).index_select(0, moba_q_sh_indices).view(1, -1)
)
dmq = torch.empty_like(moba_q)
dmkv = torch.empty_like(moba_kv)
_ = _flash_attn_varlen_backward(
dout=d_moba_output,
q=moba_q,
k=moba_kv[:, 0],
v=moba_kv[:, 1],
out=moba_output,
softmax_lse=mixed_attn_vlse,
dq=dmq,
dk=dmkv[:,0],
dv=dmkv[:,1],
cu_seqlens_q=moba_cu_seqlen_q,
cu_seqlens_k=moba_cu_seqlen_kv,
max_seqlen_q=max_seqlen,
max_seqlen_k=moba_chunk_size,
softmax_scale=softmax_scale,
causal=False,
dropout_p=0.0,
softcap=0.0,
alibi_slopes=None,
deterministic=True,
window_size_left=-1,
window_size_right=-1
)
return dq, dk, dv, None, dmq, dmkv, None, None, None, None, None
def moba_attn_varlen(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
cu_seqlens: torch.Tensor,
max_seqlen: int,
moba_chunk_size: int,
moba_topk: int,
select_mode: str = 'threshold', # "topk" or "threshold"
simsum_threshold: float = 0.25,
threshold_type: str = 'query_head',
) -> torch.Tensor:
"""
Accelerated MOBA attention for vision tasks with proper LSE normalization.
This version:
- Splits KV into chunks.
- For each query head, selects the top-k relevant KV chunks (including the self chunk)
by amplifying the diagonal (self-chunk) logits.
- Aggregates the attention outputs from the selected chunks using a log-sum-exp
reduction so that attending to each query over the selected chunks is equivalent
to the original algorithm.
"""
# Stack keys and values.
kv = torch.stack((k, v), dim=1)
seqlen, num_head, head_dim = q.shape
# Compute chunk boundaries.
cu_chunk, filtered_chunk_indices, num_filtered_chunk, chunk_to_batch = calc_chunks(
cu_seqlens, moba_chunk_size
)
self_attn_cu_seqlen = cu_chunk
# Update top-k selection to include the self chunk.
moba_topk = min(moba_topk, num_filtered_chunk)
# --- Build filtered KV from chunks ---
chunk_starts = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
chunk_ends = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
chunk_lengths = chunk_ends - chunk_starts # [num_filtered_chunk]
max_chunk_len = int(chunk_lengths.max().item())
range_tensor = torch.arange(max_chunk_len, device=kv.device, dtype=chunk_starts.dtype).unsqueeze(0)
indices = chunk_starts.unsqueeze(1) + range_tensor
indices = torch.clamp(indices, max=kv.shape[0] - 1)
valid_mask = range_tensor < chunk_lengths.unsqueeze(1)
gathered = kv[indices.view(-1)].view(num_filtered_chunk, max_chunk_len, *kv.shape[1:])
gathered = gathered * valid_mask.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1).type_as(gathered)
# Compute key_gate_weight over valid tokens.
key_values = gathered[:, :, 0].float() # [num_filtered_chunk, max_chunk_len, num_head, head_dim]
valid_mask_exp = valid_mask.unsqueeze(-1).unsqueeze(-1)
key_sum = (key_values * valid_mask_exp).sum(dim=1)
divisor = valid_mask.sum(dim=1).unsqueeze(-1).unsqueeze(-1)
key_gate_weight = key_sum / divisor # [num_filtered_chunk, num_head, head_dim]
# Compute gate logits between key_gate_weight and queries.
q_float = q.float()
# gate = torch.einsum("nhd,shd->nhs", key_gate_weight, q_float) # [num_filtered_chunk, num_head, seqlen]
gate = torch.bmm(key_gate_weight.permute(1, 0, 2), q_float.permute(1, 0, 2).transpose(1, 2)).permute(1, 0, 2)
# Amplify the diagonal (self chunk) contributions.
gate_seq_idx = torch.arange(seqlen, device=q.device, dtype=torch.int32).unsqueeze(0).expand(num_filtered_chunk, seqlen)
chunk_start = cu_chunk[filtered_chunk_indices] # [num_filtered_chunk]
chunk_end = cu_chunk[filtered_chunk_indices + 1] # [num_filtered_chunk]
gate_self_chunk_mask = ((gate_seq_idx >= chunk_start.unsqueeze(1)) &
(gate_seq_idx < chunk_end.unsqueeze(1))).unsqueeze(1).expand(-1, num_head, -1)
amplification_factor = 1e9 # Example factor; adjust as needed.
origin_gate = gate.clone()
gate = gate.clone()
if select_mode == "topk":
gate[gate_self_chunk_mask] += amplification_factor
# Exclude positions that are outside the valid batch boundaries.
batch_starts = cu_seqlens[chunk_to_batch[filtered_chunk_indices]]
batch_ends = cu_seqlens[chunk_to_batch[filtered_chunk_indices] + 1]
gate_batch_start_mask = gate_seq_idx < batch_starts.unsqueeze(1)
gate_batch_end_mask = gate_seq_idx >= batch_ends.unsqueeze(1)
gate_inf_mask = gate_batch_start_mask | gate_batch_end_mask
gate.masked_fill_(gate_inf_mask.unsqueeze(1), -float("inf"))
if select_mode == 'topk':
# We amplify self‐chunk in gate already, so self entries will rank highest.
valid_gate_mask = gate != -float("inf")
if threshold_type == 'query_head':
# === per‐<head,seq> top-k across chunks (original behavior) ===
# gate: (C, H, S)
_, gate_topk_idx = torch.topk(gate, k=moba_topk, dim=0, largest=True, sorted=False)
gate_idx_mask = torch.zeros_like(gate, dtype=torch.bool)
gate_idx_mask.scatter_(0, gate_topk_idx, True)
gate_mask = valid_gate_mask & gate_idx_mask
elif threshold_type == 'overall':
# === global top-k across all (chunk, head, seq) entries ===
C, H, S = gate.shape
flat_gate = gate.flatten()
flat_mask = valid_gate_mask.flatten()
flat_gate_masked = torch.where(flat_mask, flat_gate, -float("inf"))
# pick topk global entries
vals, idx = torch.topk(flat_gate_masked, k=moba_topk * H * S, largest=True, sorted=False)
others_mask_flat = torch.zeros_like(flat_mask, dtype=torch.bool)
others_mask_flat[idx] = True
gate_mask = (valid_gate_mask.flatten() & others_mask_flat).view(gate.shape)
elif threshold_type == 'head_global':
# per-head top-k across all chunks and sequence positions
C, H, S = gate.shape
CS = C * S
flat_gate = gate.permute(1, 0, 2).reshape(H, CS)
flat_valid = valid_gate_mask.permute(1, 0, 2).reshape(H, CS)
flat_gate_masked = torch.where(flat_valid, flat_gate, torch.full_like(flat_gate, -float('inf')))
# pick top-k indices per head
_, topk_idx = torch.topk(flat_gate_masked, k=moba_topk * S, dim=1, largest=True, sorted=False)
gate_idx_flat = torch.zeros_like(flat_valid, dtype=torch.bool)
gate_idx_flat.scatter_(1, topk_idx, True)
gate_mask = gate_idx_flat.reshape(H, C, S).permute(1, 0, 2)
else:
raise ValueError(
f"Invalid threshold_type for topk: {threshold_type}. "
"Choose 'query_head', 'block', or 'overall'."
)
elif select_mode == 'threshold':
# Delegate to the specific thresholding function
valid_gate_mask = gate != -float("inf") # (num_chunk, num_head, seqlen)
if threshold_type == 'query_head':
gate_mask = _select_threshold_query_head(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'block':
gate_mask = _select_threshold_block(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'overall':
gate_mask = _select_threshold_overall(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
elif threshold_type == 'head_global':
gate_mask = _select_threshold_head_global(gate, valid_gate_mask, gate_self_chunk_mask, simsum_threshold)
else:
raise ValueError(f"Invalid threshold_type: {threshold_type}. Choose 'query_head', 'block', or 'overall'.")
else:
raise ValueError(f"Invalid select_mode: {select_mode}. Choose 'topk' or 'threshold'.")
# eliminate self_chunk in MoBA branch
gate_mask = gate_mask & ~gate_self_chunk_mask
# if gate_mask is all false, perform flash_attn instead
if gate_mask.sum() == 0:
return flash_attn_varlen_func(
q, k, v, cu_seqlens, cu_seqlens, max_seqlen, max_seqlen, causal=False
)
# Determine which query positions are selected.
# nonzero_indices has shape [N, 3] where each row is [chunk_index, head_index, seq_index].
moba_q_indices = gate_mask.reshape(gate_mask.shape[0], -1).nonzero(as_tuple=True)[-1] # [(h s k)]
moba_q_sh_indices = (moba_q_indices % seqlen) * num_head + (moba_q_indices // seqlen)
moba_q = rearrange(q, "s h d -> (h s) d").index_select(0, moba_q_indices).unsqueeze(1)
# Build cumulative sequence lengths for the selected queries.
moba_seqlen_q = gate_mask.sum(dim=-1).flatten()
q_zero_mask = moba_seqlen_q == 0
valid_expert_mask = ~q_zero_mask
if q_zero_mask.sum() > 0:
moba_seqlen_q = moba_seqlen_q[valid_expert_mask]
moba_cu_seqlen_q = torch.cat(
(
torch.tensor([0], device=q.device, dtype=moba_seqlen_q.dtype),
moba_seqlen_q.cumsum(dim=0),
),
dim=0,
).to(torch.int32)
# Rearrange gathered KV for the MOBA branch.
experts_tensor = rearrange(gathered, "nc cl two h d -> (nc h) cl two d")
valid_expert_lengths = chunk_lengths.unsqueeze(1).expand(num_filtered_chunk, num_head).reshape(-1).to(torch.int32)
if q_zero_mask.sum() > 0:
experts_tensor = experts_tensor[valid_expert_mask]
valid_expert_lengths = valid_expert_lengths[valid_expert_mask]
seq_range = torch.arange(experts_tensor.shape[1], device=experts_tensor.device).unsqueeze(0)
mask = seq_range < valid_expert_lengths.unsqueeze(1)
moba_kv = experts_tensor[mask] # Shape: ((nc h cl_valid) two d)
moba_kv = moba_kv.unsqueeze(2) # Shape: ((nc h cl_valid) two 1 d)
moba_cu_seqlen_kv = torch.cat(
[torch.zeros(1, device=experts_tensor.device, dtype=torch.int32),
valid_expert_lengths.cumsum(dim=0)],
dim=0,
).to(torch.int32)
assert (
moba_cu_seqlen_kv.shape == moba_cu_seqlen_q.shape
), f"Mismatch between moba_cu_seqlen_kv.shape and moba_cu_seqlen_q.shape: {moba_cu_seqlen_kv.shape} vs {moba_cu_seqlen_q.shape}"
return MixedAttention.apply(
q,
k,
v,
self_attn_cu_seqlen,
moba_q,
moba_kv,
moba_cu_seqlen_q,
moba_cu_seqlen_kv,
max_seqlen,
moba_chunk_size,
moba_q_sh_indices,
)
def process_moba_input(
x,
patch_resolution,
chunk_size,
):
"""
Process inputs for the attention function.
Args:
x (torch.Tensor): Input tensor with shape [batch_size, num_patches, num_heads, head_dim].
patch_resolution (tuple): Tuple containing the patch resolution (t, h, w).
chunk_size (int): Size of the chunk. (maybe tuple or int, according to chunk type)
Returns:
torch.Tensor: Processed input tensor.
"""
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
moba_chunk_size = int(chunk_size * patch_resolution[1] * patch_resolution[2])
else:
assert isinstance(chunk_size, (Tuple, list)), f"chunk_size should be a tuple, list, or int, now it is: {type(chunk_size)}"
if len(chunk_size) == 2:
assert patch_resolution[1] % chunk_size[0] == 0 and patch_resolution[2] % chunk_size[1] == 0, f"spatial patch_resolution {patch_resolution[1:]} should be divisible by 2d chunk_size {chunk_size}"
nch, ncw = patch_resolution[1] // chunk_size[0], patch_resolution[2] // chunk_size[1]
x = rearrange(x, "b (t nch ch ncw cw) n d -> b (nch ncw t ch cw) n d", t=patch_resolution[0], nch=nch, ncw=ncw, ch=chunk_size[0], cw=chunk_size[1])
moba_chunk_size = patch_resolution[0] * chunk_size[0] * chunk_size[1]
elif len(chunk_size) == 3:
assert patch_resolution[0] % chunk_size[0] == 0 and patch_resolution[1] % chunk_size[1] == 0 and patch_resolution[2] % chunk_size[2] == 0, f"patch_resolution {patch_resolution} should be divisible by 3d chunk_size {chunk_size}"
nct, nch, ncw = patch_resolution[0] // chunk_size[0], patch_resolution[1] // chunk_size[1], patch_resolution[2] // chunk_size[2]
x = rearrange(x, "b (nct ct nch ch ncw cw) n d -> b (nct nch ncw ct ch cw) n d", nct=nct, nch=nch, ncw=ncw, ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
moba_chunk_size = chunk_size[0] * chunk_size[1] * chunk_size[2]
else:
raise ValueError(f"chunk_size should be a int, or a tuple of length 2 or 3, now it is: {len(chunk_size)}")
return x, moba_chunk_size
def process_moba_output(
x,
patch_resolution,
chunk_size,
):
if isinstance(chunk_size, float) or isinstance(chunk_size, int):
pass
elif len(chunk_size) == 2:
x = rearrange(x, "b (nch ncw t ch cw) n d -> b (t nch ch ncw cw) n d", nch=patch_resolution[1] // chunk_size[0], ncw=patch_resolution[2] // chunk_size[1], t=patch_resolution[0], ch=chunk_size[0], cw=chunk_size[1])
elif len(chunk_size) == 3:
x = rearrange(x, "b (nct nch ncw ct ch cw) n d -> b (nct ct nch ch ncw cw) n d", nct=patch_resolution[0] // chunk_size[0], nch=patch_resolution[1] // chunk_size[1], ncw=patch_resolution[2] // chunk_size[2], ct=chunk_size[0], ch=chunk_size[1], cw=chunk_size[2])
return x
# TEST
def generate_data(batch_size, seqlen, num_head, head_dim, dtype):
random.seed(0)
torch.manual_seed(0)
torch.cuda.manual_seed(0)
device = torch.cuda.current_device()
q = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
k = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
v = torch.randn((batch_size, seqlen, num_head, head_dim), requires_grad=True).to(dtype=dtype, device='cuda')
print(f"q.shape: {q.shape}, k.shape: {k.shape}, v.shape: {v.shape}")
cu_seqlens = torch.arange(0, q.shape[0] * q.shape[1] + 1, q.shape[1], dtype=torch.int32, device='cuda')
max_seqlen = q.shape[1]
q = rearrange(q, "b s ... -> (b s) ...")
k = rearrange(k, "b s ... -> (b s) ...")
v = rearrange(v, "b s ... -> (b s) ...")
return q, k, v, cu_seqlens, max_seqlen
def test_attn_varlen_moba_speed(batch, head, seqlen, head_dim, moba_chunk_size, moba_topk, dtype=torch.bfloat16, select_mode='threshold', simsum_threshold=0.25, threshold_type='query_head'):
"""Speed test comparing flash_attn vs moba_attention"""
# Get data
q, k, v, cu_seqlen, max_seqlen = generate_data(batch, seqlen, head, head_dim, dtype)
print(f"batch:{batch} head:{head} seqlen:{seqlen} chunk:{moba_chunk_size} topk:{moba_topk} select_mode: {select_mode} simsum_threshold:{simsum_threshold}")
vo_grad = torch.randn_like(q)
# Warmup
warmup_iters = 3
perf_test_iters = 10
# Warmup
for _ in range(warmup_iters):
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
torch.autograd.backward(o, vo_grad)
torch.cuda.synchronize()
start_flash = time.perf_counter()
for _ in range(perf_test_iters):
o = flash_attn_varlen_func(q, k, v, cu_seqlen, cu_seqlen, max_seqlen, max_seqlen, causal=False)
torch.autograd.backward(o, vo_grad)
torch.cuda.synchronize()
time_flash = (time.perf_counter() - start_flash) / perf_test_iters * 1000
# Warmup
for _ in range(warmup_iters):
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
torch.autograd.backward(om, vo_grad)
torch.cuda.synchronize()
start_moba = time.perf_counter()
for _ in range(perf_test_iters):
om = moba_attn_varlen(q, k, v, cu_seqlen, max_seqlen, moba_chunk_size=moba_chunk_size, moba_topk=moba_topk, select_mode=select_mode, simsum_threshold=simsum_threshold, threshold_type=threshold_type)
torch.autograd.backward(om, vo_grad)
torch.cuda.synchronize()
time_moba = (time.perf_counter() - start_moba) / perf_test_iters * 1000
print(f"Flash: {time_flash:.2f}ms, MoBA: {time_moba:.2f}ms")
print(f"Speedup: {time_flash / time_moba:.2f}x")
if __name__ == "__main__":
"""
CUDA_VISIBLE_DEVICES=1 \
python -u csrc/attn/vmoba_attn/vmoba/vmoba.py
"""
test_attn_varlen_moba_speed(batch=1, head=12, seqlen=32760, head_dim=128, moba_chunk_size=32760 // 3 // 6 // 4, moba_topk=3, select_mode='threshold', simsum_threshold=0.3, threshold_type='query_head')
@@ -98,7 +98,6 @@ dmd_args=(
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
--master_port $MASTER_PORT \
fastvideo/training/wan_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
@@ -1,112 +0,0 @@
#!/bin/bash
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export MASTER_PORT=29501
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
# Configs
NUM_GPUS=1
MODEL_PATH="Wan-AI/Wan2.2-TI2V-5B-Diffusers"
DATA_DIR="data/crush-smol_processed_ti2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/distill/Wan2.2-TI2V-5B-Diffusers/crush_smol/validation.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name wan_t2v_distill_dmd_VSA
--output_dir="checkpoints/wan_t2v_finetune"
--max_train_steps=4000
--train_batch_size=1
--train_sp_batch_size 1
--gradient_accumulation_steps=1
--num_latent_t 31
--num_height 704
--num_width 1280
--num_frames 121
--enable_gradient_checkpointing_type "full"
--training_state_checkpointing_steps=500
--weight_only_checkpointing_steps=500
--lora_rank 32
--lora_training True
)
# Parallel arguments
parallel_args=(
--num_gpus 1
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim 1
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 4
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 200
--validation_sampling_steps "3"
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate=1e-4
--mixed_precision="bf16"
--weight_decay 0.01
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--ema_start_step 0
--flow_shift 8
--seed 1000
)
# DMD arguments
dmd_args=(
--dmd_denoising_steps '1000,757,522'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--generator_update_interval 5
--real_score_guidance_scale 3.5
--VSA_sparsity 0.8
)
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
--master_port $MASTER_PORT \
fastvideo/training/wan_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
@@ -1,47 +0,0 @@
A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.
The video shows a green and orange object being flattened as if it were under a hydraulic press, with the press moving down and compressing the object.
The video shows a cylindrical object with a cityscape image being flattened as if it were under a hydraulic press. The object is placed on a metal platform, and a large, striped cylinder presses down on it, causing it to collapse and release a liquid inside. The background features a green wall with a yellow and red warning sign.
A red toy car is being crushed by a large hydraulic press, which is flattening objects as if they were under a hydraulic press.
A large, cylindrical object is seen pressing down on a small orange ball, causing it to flatten as if it were under a hydraulic press. The background features a green wall with yellow and red warning signs.
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is shown compressing a wooden object, which shatters into small pieces. The background features a green wall with a yellow sign displaying a lightning bolt.
A large metal cylinder is seen descending, flattening objects as if they were under a hydraulic press. The cylinder compresses a stack of matches and boxes, causing them to crumble into small pieces. The scene is set against a green background with yellow and red signs.
A large metal press is shown compressing a pile of colorful macarons, flattening them as if they were under a hydraulic press. The press moves down, crushing the macarons into a pile of crumbs and squishing the colorful filling out.
The video shows a metal press flattening objects as if they were under a hydraulic press. The press is pressing down on a pile of colorful gummy candies, squishing them into a pile of squiggly shapes. The press is made of metal and has a large base, and the gummy candies are of various colors, including red, green, and orange. The background is a green wall, and the press is placed on a metal surface.
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
The video shows a stack of colorful sponges being flattened as if they were under a hydraulic press. The sponges, which are pink, white, blue, and green, are compressed into a smaller size, demonstrating the press's power. The background features a green wall with a yellow and red sign, adding context to the setting.
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, leaving a pile of debris around it.
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press.
The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.
The video shows a close-up of a metal cylinder pressing down on a yellow object, which is being flattened as if it were under a hydraulic press. The cylinder is positioned above the object, and the force is causing the object to compress and spread out, creating a visible deformation. The background is blurred, focusing attention on the action of the cylinder and the object being flattened.
A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.
The video shows a hydraulic press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing two colorful objects that resemble sandwiches. The press is yellow and black striped, and the objects being flattened are placed on a metal plate. The background is green, and the press is moving down, compressing the objects.
The scene shows a metal press with a yellow and black striped pattern, holding a container filled with chocolate. A metal cylinder is descending, flattening the chocolate as if it were under a hydraulic press. The background is a green wall, and the press is mounted on a sturdy metal frame.
The video shows a colorful sponge being flattened as if it were under a hydraulic press, with the sponge being compressed and eventually flattened into a thin layer.
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is pressing down on a stack of wooden blocks, causing them to crumble and break apart. The press is black and yellow striped, and the wooden blocks are small and rectangular. The background is green, and the press is sitting on a metal table.
A pile of colorful candies is being flattened by a hydraulic press, causing them to crumble into small pieces.
The video shows a stack of colorful sponges being flattened by a large, cylindrical object, which appears to be a hydraulic press. The sponges, which are pink, blue, white, and green, are compressed into a single layer, demonstrating the press's powerful force. The background features a green wall with a yellow and red sign, adding context to the industrial setting.
A bowling ball is placed on a metal platform, and a large metal cylinder descends from above, flattening the ball as if it were under a hydraulic press. The ball is crushed into a flat, round shape, demonstrating the immense pressure applied by the cylinder.
A large metal cylinder with yellow and black stripes is seen pressing down on a pile of popcorn, flattening the objects as if they were under a hydraulic press. The popcorn is crushed and scattered around the base of the cylinder, creating a satisfying visual effect.
The video shows a hydraulic press in action, flattening objects as if they were under a hydraulic press. The press is composed of a large, cylindrical metal cylinder with yellow and black stripes, and a metal base. The objects being flattened are two cylindrical blocks of cotton candy, one pink and one blue. The press is positioned on a metal table, and the background features a green wall with a yellow and red sign.
The video shows a large orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.
The video shows a cylindrical object being pressed down onto a flat surface, causing the objects beneath it to be flattened as if they were under a hydraulic press. The objects being flattened appear to be yellow and are being crushed into a pile of debris. The background is a greenish-gray color, and the surface on which the objects are being flattened is metallic and shiny.
A green and blue object with a spiky texture is being flattened by a large, cylindrical metal press, demonstrating its resilience and durability.
The video shows a stack of caramelized sugar cubes being flattened as if they were under a hydraulic press, resulting in a messy pile of broken sugar on the table.
A large metal cylinder is seen pressing down on a pile of colorful jelly beans, flattening them as if they were under a hydraulic press.
The video shows a machine with a yellow and black striped cylinder pressing down on a stack of colorful sponges, flattening them as if they were under a hydraulic press. The machine is situated in a green-walled room with warning signs in the background.
The video shows a machine with a yellow and black striped cylinder, which is pressing down on two colorful objects, flattening them as if they were under a hydraulic press. The machine appears to be in a workshop or industrial setting, with a green wall in the background. The objects being flattened are green and orange, and the machine is covered in dirt and grime, indicating it has been used frequently.
The video shows a large, industrial press flattening objects as if they were under a hydraulic press. The press is shown in action, compressing a pile of pink objects into a pile of crumbs. The press is large and metallic, with a yellow and black striped pattern on its side. The background is a green wall with a yellow warning sign.
The video shows a pink, sparkly ball being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and segments.
The video shows a machine with a yellow and black striped cylinder, which is flattening objects as if they were under a hydraulic press. The machine is pressing down on two colorful objects, causing them to compress and flatten. The background is a green wall, and the machine appears to be in a workshop or industrial setting.
The video shows a large, yellow and black striped cylinder flattening objects as if they were under a hydraulic press. The objects being flattened are pink and are being crushed into small pieces. The background is a green wall with a yellow sign.
The video shows a machine with a yellow and black striped cylinder pressing down on two colorful objects, which are flattened as if they were under a hydraulic press. The machine is positioned on a metal platform, and the background is a green wall.
A green cube is being compressed by a hydraulic press, which flattens the object as if it were under a hydraulic press. The press is shown in action, with the cube being squeezed into a smaller shape.
A pink, sparkly ball is being crushed by a large, rusty cylinder, which flattens the objects as if they were under a hydraulic press.
A red cabbage is being crushed by a hydraulic press, which flattens the objects as if they were under a hydraulic press. The press is shown in action, compressing the cabbage into a smaller, more compact form.
A lime is being crushed by a hydraulic press, causing it to flatten and burst open, releasing its juice and pulp.
A large metal press is shown compressing a stack of burgers, causing them to be flattened and crushed into a pile of ground meat.
A pizza is being crushed by a hydraulic press, causing the toppings to spread out and the crust to crumble.
A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.
A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.
A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.
@@ -1,93 +0,0 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_crush_smol"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "wan_ode_init_crush_smol"
--max_train_steps 6000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 21
--num_height 480
--num_width 832
--num_frames 77
--warp_denoising_step
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim 1
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 6e-6
--mixed_precision "bf16"
--checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -1,24 +0,0 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="$(dirname "$0")/crush_smol_prompts.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 1 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "ode_trajectory"
@@ -1,40 +0,0 @@
{
"data": [
{
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A stylish woman walks down a Tokyo street filled with warm glowing neon and animated city signage. She wears a black leather jacket, a long red dress, and black boots, and carries a black purse. She wears sunglasses and red lipstick. She walks confidently and casually. The street is damp and reflective, creating a mirror effect of the colorful lights. Many pedestrians walk about.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A white and orange tabby cat is seen happily darting through a dense garden, as if chasing something. Its eyes are wide and happy as it jogs forward, scanning the branches, flowers, and leaves as it walks. The path is narrow as it makes its way between all the plants. the scene is captured from a ground-level angle, following the cat closely, giving a low and intimate perspective. The image is cinematic with warm tones and a grainy texture. The scattered daylight between the leaves and plants above creates a warm contrast, accentuating the cat’s orange fur. The shot is clear and sharp, with a shallow depth of field.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
"image_path": null,
"video_path": null,
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -28,4 +28,4 @@
"num_frames": 77
}
]
}
}
-214
View File
@@ -1,214 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import re
from dataclasses import dataclass
import torch
from einops import rearrange
from flash_attn.bert_padding import pad_input
from csrc.attn.vmoba_attn.vmoba import (moba_attn_varlen, process_moba_input,
process_moba_output)
from fastvideo.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.logger import init_logger
logger = init_logger(__name__)
class VMOBAAttentionBackend(AttentionBackend):
accept_output_buffer: bool = True
@staticmethod
def get_name() -> str:
return "VMOBA_ATTN"
@staticmethod
def get_impl_cls() -> type["VMOBAAttentionImpl"]:
return VMOBAAttentionImpl
@staticmethod
def get_metadata_cls() -> type["VideoMobaAttentionMetadata"]:
return VideoMobaAttentionMetadata
@staticmethod
def get_builder_cls() -> type["VideoMobaAttentionMetadataBuilder"]:
return VideoMobaAttentionMetadataBuilder
@dataclass
class VideoMobaAttentionMetadata(AttentionMetadata):
current_timestep: int
temporal_chunk_size: int
temporal_topk: int
spatial_chunk_size: tuple[int, int]
spatial_topk: int
st_chunk_size: tuple[int, int, int]
st_topk: int
moba_select_mode: str
moba_threshold: float
moba_threshold_type: str
patch_resolution: list[int]
first_full_step: int = 12
first_full_layer: int = 0
# temporal_layer -> spatial_layer -> st_layer
temporal_layer: int = 1
spatial_layer: int = 1
st_layer: int = 1
class VideoMobaAttentionMetadataBuilder(AttentionMetadataBuilder):
def __init__(self):
pass
def prepare(self):
pass
def build( # type: ignore
self,
current_timestep: int,
raw_latent_shape: tuple[int, int, int],
patch_size: tuple[int, int, int],
temporal_chunk_size: int,
temporal_topk: int,
spatial_chunk_size: tuple[int, int],
spatial_topk: int,
st_chunk_size: tuple[int, int, int],
st_topk: int,
moba_select_mode: str = 'threshold',
moba_threshold: float = 0.25,
moba_threshold_type: str = 'query_head',
device: torch.device = None,
first_full_layer: int = 0,
first_full_step: int = 12,
temporal_layer: int = 1,
spatial_layer: int = 1,
st_layer: int = 1,
**kwargs,
) -> VideoMobaAttentionMetadata:
if device is None:
device = torch.device("cpu")
assert raw_latent_shape[0] % patch_size[0] == 0 and raw_latent_shape[
1] % patch_size[1] == 0 and raw_latent_shape[2] % patch_size[
2] == 0, f"spatial patch_resolution {raw_latent_shape} should be divisible by patch_size {patch_size}"
patch_resolution = [
t // pt for t, pt in zip(raw_latent_shape, patch_size, strict=False)
]
return VideoMobaAttentionMetadata(
current_timestep=current_timestep,
temporal_chunk_size=temporal_chunk_size,
temporal_topk=temporal_topk,
spatial_chunk_size=spatial_chunk_size,
spatial_topk=spatial_topk,
st_chunk_size=st_chunk_size,
st_topk=st_topk,
moba_select_mode=moba_select_mode,
moba_threshold=moba_threshold,
moba_threshold_type=moba_threshold_type,
patch_resolution=patch_resolution,
first_full_layer=first_full_layer,
first_full_step=first_full_step,
temporal_layer=temporal_layer,
spatial_layer=spatial_layer,
st_layer=st_layer,
)
class VMOBAAttentionImpl(AttentionImpl):
def __init__(self,
num_heads,
head_size,
softmax_scale,
causal=False,
num_kv_heads=None,
prefix="",
**extra_impl_args) -> None:
self.prefix = prefix
self.layer_idx = self._get_layer_idx(prefix)
def _get_layer_idx(self, prefix: str) -> int | None:
match = re.search(r"blocks\.(\d+)", prefix)
if not match:
raise ValueError(f"Invalid prefix: {prefix}")
return int(match.group(1))
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: AttentionMetadata,
) -> torch.Tensor:
"""
query: [B, L, H, D]
key: [B, L, H, D]
value: [B, L, H, D]
attn_metadata: AttentionMetadata
"""
batch_size, sequence_length, num_heads, head_dim = query.shape
# select chunk type according to layer idx:
loop_layer_num = attn_metadata.temporal_layer + attn_metadata.spatial_layer + attn_metadata.st_layer
moba_layer = self.layer_idx - attn_metadata.first_full_layer
if moba_layer % loop_layer_num < attn_metadata.temporal_layer:
moba_chunk_size = attn_metadata.temporal_chunk_size
moba_topk = attn_metadata.temporal_topk
elif moba_layer % loop_layer_num < attn_metadata.temporal_layer + attn_metadata.spatial_layer:
moba_chunk_size = attn_metadata.spatial_chunk_size
moba_topk = attn_metadata.spatial_topk
elif moba_layer % loop_layer_num < attn_metadata.temporal_layer + attn_metadata.spatial_layer + attn_metadata.st_layer:
moba_chunk_size = attn_metadata.st_chunk_size
moba_topk = attn_metadata.st_topk
# torch.distributed.breakpoint()
query, chunk_size = process_moba_input(query,
attn_metadata.patch_resolution,
moba_chunk_size)
key, chunk_size = process_moba_input(key,
attn_metadata.patch_resolution,
moba_chunk_size)
value, chunk_size = process_moba_input(value,
attn_metadata.patch_resolution,
moba_chunk_size)
max_seqlen = query.shape[1]
indices_q = torch.arange(0,
query.shape[0] * query.shape[1],
device=query.device)
cu_seqlens = torch.arange(0,
query.shape[0] * query.shape[1] + 1,
query.shape[1],
dtype=torch.int32,
device=query.device)
query = rearrange(query, "b s ... -> (b s) ...")
key = rearrange(key, "b s ... -> (b s) ...")
value = rearrange(value, "b s ... -> (b s) ...")
# current_timestep=attn_metadata.current_timestep
hidden_states = moba_attn_varlen(
query,
key,
value,
cu_seqlens=cu_seqlens,
max_seqlen=max_seqlen,
moba_chunk_size=chunk_size,
moba_topk=moba_topk,
select_mode=attn_metadata.moba_select_mode,
simsum_threshold=attn_metadata.moba_threshold,
threshold_type=attn_metadata.moba_threshold_type,
)
hidden_states = pad_input(hidden_states, indices_q, batch_size,
sequence_length)
hidden_states = process_moba_output(hidden_states,
attn_metadata.patch_resolution,
moba_chunk_size)
return hidden_states
@@ -1,16 +0,0 @@
{
"temporal_chunk_size": 2,
"temporal_topk": 2,
"spatial_chunk_size": [4, 13],
"spatial_topk": 6,
"st_chunk_size": [4, 4, 13],
"st_topk": 18,
"moba_select_mode": "topk",
"moba_threshold": 0.25,
"moba_threshold_type": "query_head",
"first_full_layer": 0,
"first_full_step": 12,
"temporal_layer": 1,
"spatial_layer": 1,
"st_layer": 1
}
@@ -1,16 +0,0 @@
{
"temporal_chunk_size": 2,
"temporal_topk": 3,
"spatial_chunk_size": [3, 4],
"spatial_topk": 20,
"st_chunk_size": [4, 6, 4],
"st_topk": 15,
"moba_select_mode": "threshold",
"moba_threshold": 0.25,
"moba_threshold_type": "query_head",
"first_full_layer": 0,
"first_full_step": 12,
"temporal_layer": 1,
"spatial_layer": 1,
"st_layer": 1
}
+3 -7
View File
@@ -15,13 +15,9 @@ class DiTArchConfig(ArchConfig):
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
_supported_attention_backends: tuple[AttentionBackendEnum, ...] = (
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.VMOBA_ATTN,
)
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
+2 -3
View File
@@ -4,13 +4,12 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
WanI2V480PConfig, WanI2V720PConfig,
from fastvideo.configs.pipelines.wan import (WanI2V480PConfig, WanI2V720PConfig,
WanT2V480PConfig, WanT2V720PConfig)
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
"SelfForcingWanT2V480PConfig", "get_pipeline_config_cls_from_name"
"get_pipeline_config_cls_from_name"
]
+3 -5
View File
@@ -11,9 +11,9 @@ from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
# isort: off
from fastvideo.configs.pipelines.wan import (
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
SelfForcingWanT2V480PConfig)
SelfForcingWanT2V480PConfig, Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config,
Wan2_2_TI2V_5B_Config, WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig,
WanT2V720PConfig)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
@@ -48,7 +48,6 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
"stepvideo": lambda id: "stepvideo" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -61,7 +60,6 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V480PConfig,
"wandmdpipeline": FastWan2_1_T2V_480P_Config,
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
"stepvideo": StepVideoT2VConfig
# Other fallbacks by architecture
}
+11 -13
View File
@@ -12,13 +12,13 @@ from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
mask: torch.Tensor = outputs.attention_mask
hidden_state: torch.Tensor = outputs.last_hidden_state
def t5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
mask: torch.tensor = outputs.attention_mask
hidden_state: torch.tensor = outputs.last_hidden_state
seq_lens = mask.gt(0).sum(dim=1).long()
assert torch.isnan(hidden_state).sum() == 0
prompt_embeds = [u[:v] for u, v in zip(hidden_state, seq_lens, strict=True)]
prompt_embeds_tensor: torch.Tensor = torch.stack([
prompt_embeds_tensor: torch.tensor = torch.stack([
torch.cat([u, u.new_zeros(512 - u.size(0), u.size(1))])
for u in prompt_embeds
],
@@ -39,12 +39,12 @@ class WanT2V480PConfig(PipelineConfig):
vae_sp: bool = False
# Denoising stage
flow_shift: float | None = 3.0
flow_shift: int = 3
# Text encoding stage
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (T5Config(), ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.tensor],
...] = field(default_factory=lambda:
(t5_postprocess_text, ))
@@ -68,7 +68,7 @@ class WanT2V720PConfig(WanT2V480PConfig):
# WanConfig-specific parameters with defaults
# Denoising stage
flow_shift: float | None = 5.0
flow_shift: int = 5
@dataclass
@@ -94,7 +94,7 @@ class WanI2V720PConfig(WanI2V480PConfig):
# WanConfig-specific parameters with defaults
# Denoising stage
flow_shift: float | None = 5.0
flow_shift: int = 5
@dataclass
@@ -104,7 +104,7 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
# WanConfig-specific parameters with defaults
# Denoising stage
flow_shift: float | None = 8.0
flow_shift: int = 8
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 757, 522])
@@ -115,7 +115,7 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
@dataclass
class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
flow_shift: float | None = 5.0
flow_shift: int = 5
ti2v_task: bool = True
def __post_init__(self) -> None:
@@ -125,7 +125,7 @@ class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
@dataclass
class FastWan2_2_TI2V_5B_Config(Wan2_2_TI2V_5B_Config):
flow_shift: float | None = 5.0
flow_shift: int = 5
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 757, 522])
@@ -146,7 +146,5 @@ class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
@dataclass
class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
is_causal: bool = True
flow_shift: float | None = 5.0
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 750, 500, 250])
warp_denoising_step: bool = True
-21
View File
@@ -47,8 +47,6 @@ class SamplingParam:
# Misc
save_video: bool = True
return_frames: bool = False
return_trajectory_latents: bool = False # returns all latents for each timestep
return_trajectory_decoded: bool = False # returns decoded latents for each timestep
def __post_init__(self) -> None:
self.data_type = "video" if self.num_frames > 1 else "image"
@@ -193,25 +191,6 @@ class SamplingParam:
default=SamplingParam.image_path,
help="Path to input image for image-to-video generation",
)
parser.add_argument(
"--moba-config-path",
type=str,
default=None,
help=
"Path to a JSON file containing V-MoBA specific configurations.",
)
parser.add_argument(
"--return-trajectory-latents",
action="store_true",
default=SamplingParam.return_trajectory_latents,
help="Whether to return the trajectory",
)
parser.add_argument(
"--return-trajectory-decoded",
action="store_true",
default=SamplingParam.return_trajectory_decoded,
help="Whether to return the decoded trajectory",
)
return parser
+2 -2
View File
@@ -37,7 +37,7 @@ def getdataset(args) -> VideoCaptionMergedDataset:
temporal_sample=temporal_sample,
transform_topcrop=transform_topcrop,
seed=args.seed)
def gettextdataset(args) -> TextDataset:
return TextDataset(data_merge_path=args.data_merge_path,
@@ -48,4 +48,4 @@ def gettextdataset(args) -> TextDataset:
__all__ = [
"build_parquet_map_style_dataloader", "ValidationDataset",
"VideoCaptionMergedDataset", "TextDataset"
]
]
-264
View File
@@ -1,264 +0,0 @@
"""
Utilities for converting preprocessing records (dicts) into Arrow tables and
writing Parquet datasets in fixed-size chunks.
This module centralizes table construction and Parquet file writing so
pipelines only need to define their PyArrow schema and produce per-sample
record dictionaries.
Key APIs:
- records_to_table(records, schema): Safely convert a list of dictionaries into
a pa.Table, casting to the provided schema.
- ParquetDatasetWriter: Buffer tables and flush to a directory as multiple
Parquet files with a fixed number of rows per file. Uses temporary files and
atomic rename to avoid partially written outputs.
"""
from __future__ import annotations
import multiprocessing
import os
from concurrent.futures import ProcessPoolExecutor
from typing import Any
import pyarrow as pa
import pyarrow.parquet as pq
def records_to_table(records: list[dict[str, Any]], schema: pa.Schema) -> pa.Table:
"""Build a PyArrow table from Python record dicts using an explicit schema.
Arrow will cast values to the target schema when possible (e.g., promoting
Python ints/floats to pa.int64/pa.float64), eliminating hand-written per-
field array construction.
Args:
records: List of dictionaries, each representing one row. Keys must
match schema field names.
schema: Target PyArrow schema. Controls field names and types.
Returns:
pa.Table: In-memory table matching the provided schema. If ``records``
is empty, returns an empty table with the given schema.
"""
if not records:
return pa.table({}, schema=schema)
return pa.Table.from_pylist(records, schema=schema)
class ParquetDatasetWriter:
"""Accumulate tables and flush them to a Parquet directory in fixed-size chunks.
Behavior:
- Writes files under worker-specific subdirectories for parallelism.
- Uses temporary files and atomic rename to avoid partial files being left
behind on failure.
- Only full chunks of ``samples_per_file`` rows are written on each flush;
any remainder rows are re-buffered for the next flush.
Note:
- Instances are not meant to be shared across processes. Create one writer
per process if using multiprocessing.
"""
def __init__(self, out_dir: str, samples_per_file: int, compression: str = "zstd") -> None:
"""Initialize the dataset writer.
Args:
out_dir: Output directory where Parquet files will be written.
samples_per_file: Fixed number of rows per Parquet file.
compression: Compression codec passed to ``pyarrow.parquet.write_table``
(e.g., ``"zstd"``, ``"snappy"``, ``"gzip"``).
"""
self.out_dir = out_dir
self.samples_per_file = max(int(samples_per_file), 1)
self.compression = compression
os.makedirs(self.out_dir, exist_ok=True)
self._tables: list[pa.Table] = []
def append_table(self, table: pa.Table) -> None:
"""Append a non-empty table to the internal buffer.
Args:
table: A ``pa.Table`` to buffer. Empty or ``None`` tables are ignored.
"""
if table is None or len(table) == 0:
return
self._tables.append(table)
def _combine(self) -> pa.Table | None:
"""Combine all buffered tables into a single table, if any.
Returns:
A concatenated table, a single table if only one was buffered, or
``None`` if no tables are buffered.
"""
if not self._tables:
return None
if len(self._tables) == 1:
return self._tables[0]
return pa.concat_tables(self._tables, promote_options='none')
def flush(self, num_workers: int | None = None, write_remainder: bool = False) -> int:
"""Write accumulated tables to disk and clear the written portion.
Only complete chunks of size ``samples_per_file`` are written. Any
remainder rows are kept buffered for the next flush.
Args:
num_workers: Optional override for the number of parallel workers
used to write chunks. Defaults to ``min(cpu_count, chunks)``.
write_remainder: If True, also write any leftover rows (< samples_per_file)
as a final small Parquet file (useful for the last flush at the
end of preprocessing).
Returns:
int: Number of rows successfully written in this flush call.
"""
combined = self._combine()
self._tables = []
if combined is None or len(combined) == 0:
return 0
num_samples = len(combined)
total_chunks = num_samples // self.samples_per_file
if total_chunks == 0:
if not write_remainder:
# Not enough to form a full chunk; keep buffered for next round
# Re-buffer and return 0 written
self._tables = [combined]
return 0
# Last flush: write the small remainder as a final file in worker_0
worker_dir = os.path.join(self.out_dir, "worker_0")
os.makedirs(worker_dir, exist_ok=True)
# Determine next index
num_parquets = 0
for _, _, files in os.walk(worker_dir):
for file in files:
if file.endswith('.parquet'):
num_parquets += 1
chunk_path = os.path.join(worker_dir, f"data_chunk_{num_parquets}.parquet")
temp_path = chunk_path + '.tmp'
pq.write_table(combined, temp_path, compression=self.compression)
if os.path.exists(chunk_path):
os.remove(chunk_path)
os.rename(temp_path, chunk_path)
return num_samples
# Only write full chunks; keep remainder for next flush
written_rows = total_chunks * self.samples_per_file
remainder = num_samples - written_rows
table_to_write = combined.slice(0, written_rows)
remainder_table = combined.slice(written_rows, remainder) if remainder > 0 else None
if remainder_table is not None and len(remainder_table) > 0:
if write_remainder:
# Write the remainder as a final small file (worker_0)
worker_dir = os.path.join(self.out_dir, "worker_0")
os.makedirs(worker_dir, exist_ok=True)
num_parquets = 0
for _, _, files in os.walk(worker_dir):
for file in files:
if file.endswith('.parquet'):
num_parquets += 1
remainder_path = os.path.join(worker_dir,
f"data_chunk_{num_parquets}.parquet")
temp_path = remainder_path + '.tmp'
pq.write_table(remainder_table,
temp_path,
compression=self.compression)
if os.path.exists(remainder_path):
os.remove(remainder_path)
os.rename(temp_path, remainder_path)
else:
self._tables = [remainder_table]
# Parallel write by chunk ranges
if num_workers is None:
num_workers = min(multiprocessing.cpu_count(), max(total_chunks, 1))
num_workers = max(int(num_workers), 1)
chunks_per_worker = (total_chunks + num_workers - 1) // num_workers
work_ranges: list[tuple[int, int, pa.Table, int, str, int, str]] = []
for worker_id in range(num_workers):
start_chunk = worker_id * chunks_per_worker
end_chunk = min((worker_id + 1) * chunks_per_worker, total_chunks)
if start_chunk < end_chunk:
work_ranges.append(
(
start_chunk,
end_chunk,
table_to_write,
worker_id,
self.out_dir,
self.samples_per_file,
self.compression,
)
)
written_total = 0
if len(work_ranges) == 1:
written_total += _process_chunk_range(work_ranges[0])
return written_total
with ProcessPoolExecutor(max_workers=num_workers) as executor:
futures = [executor.submit(_process_chunk_range, args) for args in work_ranges]
for f in futures:
written_total += f.result()
return written_total + (len(remainder_table) if write_remainder and remainder_table is not None else 0)
def _process_chunk_range(args: Any) -> int:
"""Worker function to write a contiguous range of chunk files.
Args:
args: Tuple containing
- start_chunk (int): inclusive start chunk index
- end_chunk (int): exclusive end chunk index
- table (pa.Table): concatenated table containing all rows to write
- worker_id (int): numeric worker identifier
- output_dir (str): base output directory
- samples_per_file (int): rows per chunk file
- compression (str): compression codec for Parquet
Returns:
int: Total number of rows written by this worker.
"""
start_chunk, end_chunk, table, worker_id, output_dir, samples_per_file, compression = args
total_written = 0
num_samples = len(table)
worker_dir = os.path.join(output_dir, f"worker_{worker_id}")
os.makedirs(worker_dir, exist_ok=True)
# Offset to continue numbering if files exist
num_parquets = 0
for root, _, files in os.walk(worker_dir):
for file in files:
if file.endswith('.parquet'):
num_parquets += 1
for i in range(start_chunk, end_chunk):
start_sample = i * samples_per_file
end_sample = min((i + 1) * samples_per_file, num_samples)
if end_sample <= start_sample:
continue
chunk = table.slice(start_sample, end_sample - start_sample)
chunk_path = os.path.join(worker_dir, f"data_chunk_{i + num_parquets}.parquet")
temp_path = chunk_path + '.tmp'
try:
pq.write_table(chunk, temp_path, compression=compression)
if os.path.exists(chunk_path):
os.remove(chunk_path)
os.rename(temp_path, chunk_path)
total_written += len(chunk)
except Exception:
if os.path.exists(temp_path):
os.remove(temp_path)
raise
return total_written
+50 -1
View File
@@ -50,7 +50,6 @@ pyarrow_schema_i2v = pa.schema([
pa.field("fps", pa.float64()),
])
pyarrow_schema_t2v = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
@@ -80,6 +79,45 @@ pyarrow_schema_t2v = pa.schema([
pa.field("fps", pa.float64()),
])
pyarrow_schema_ode_trajectory = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("vae_latent_bytes", pa.binary()),
# e.g., [C, T, H, W] or [C, H, W]
pa.field("vae_latent_shape", pa.list_(pa.int64())),
# e.g., 'float32'
pa.field("vae_latent_dtype", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
# e.g., [SeqLen, Dim]
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
# I2V
pa.field("image_condition_latents_bytes", pa.binary()),
pa.field("image_condition_latents_shape", pa.list_(pa.int64())),
pa.field("image_condition_latents_dtype", pa.string()),
# --- ODE Trajectory ---
pa.field("trajectory_latents_bytes", pa.binary()),
pa.field("trajectory_latents_shape", pa.list_(pa.int64())),
pa.field("trajectory_latents_dtype", pa.string()),
pa.field("trajectory_timesteps_bytes", pa.binary()),
pa.field("trajectory_timesteps_shape", pa.list_(pa.int64())),
pa.field("trajectory_timesteps_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
pyarrow_schema_ode_trajectory_text_only = pa.schema([
pa.field("id", pa.string()),
@@ -102,3 +140,14 @@ pyarrow_schema_ode_trajectory_text_only = pa.schema([
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # Always 'text' for text-only
])
pyarrow_schema_text_only = pa.schema([
pa.field("id", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
# e.g., [SeqLen, Dim]
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
])
+1 -1
View File
@@ -758,4 +758,4 @@ class TextDataset(torch.utils.data.IterableDataset,
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
"""Load state dict from checkpoint."""
self.processed_batches = state_dict["processed_batches"]
self.processed_batches = state_dict["processed_batches"]
+1 -1
View File
@@ -5,7 +5,7 @@ import numpy as np
import torch
def pad(t: torch.Tensor, padding_length: int) -> tuple[torch.Tensor, torch.Tensor]:
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
"""
Pad or crop an embedding [L, D] to exactly padding_length tokens.
Return:
-3
View File
@@ -344,9 +344,6 @@ class VideoGenerator:
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time,
"logging_info": logging_info,
"trajectory": output_batch.trajectory_latents,
"trajectory_timesteps": output_batch.trajectory_timesteps,
"trajectory_decoded": output_batch.trajectory_decoded,
}
def set_lora_adapter(self,
+1 -24
View File
@@ -1,9 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
# Inspired by SGLang: https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/server_args.py
"""The arguments of FastVideo Inference."""
import argparse
import dataclasses
import json
from contextlib import contextmanager
from dataclasses import field
from enum import Enum
@@ -139,10 +139,6 @@ class FastVideoArgs:
# VSA parameters
VSA_sparsity: float = 0.0 # inference/validation sparsity
# V-MoBA parameters
moba_config_path: str | None = None
moba_config: dict[str, Any] = field(default_factory=dict)
# Master port for distributed training/inference
master_port: int | None = None
@@ -170,16 +166,6 @@ class FastVideoArgs:
return not self.inference_mode
def __post_init__(self):
if self.moba_config_path:
try:
with open(self.moba_config_path) as f:
self.moba_config = json.load(f)
logger.info("Loaded V-MoBA config from %s",
self.moba_config_path)
except (FileNotFoundError, json.JSONDecodeError) as e:
logger.error("Failed to load V-MoBA config from %s: %s",
self.moba_config_path, e)
raise
self.check_fastvideo_args()
@staticmethod
@@ -999,15 +985,6 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--lora-rank", type=int, help="LoRA rank")
parser.add_argument("--lora-alpha", type=int, help="LoRA alpha")
# V-MoBA parameters
parser.add_argument(
"--moba-config-path",
type=str,
default=None,
help=
"Path to a JSON file containing V-MoBA specific configurations.",
)
# Distillation arguments
parser.add_argument("--generator-update-interval",
type=int,
+6 -60
View File
@@ -100,16 +100,7 @@ class ScaleResidual(nn.Module):
def forward(self, residual: torch.Tensor, x: torch.Tensor,
gate: torch.Tensor) -> torch.Tensor:
"""Apply gated residual connection."""
# x.shape: [batch_size, seq_len, inner_dim]
if gate.dim() == 4:
# gate.shape: [batch_size, num_frames, 1, inner_dim]
num_frames = gate.shape[1]
frame_seqlen = x.shape[1] // num_frames
return residual + (x.unflatten(
dim=1, sizes=(num_frames, frame_seqlen)) * gate).flatten(1, 2)
else:
# gate.shape: [batch_size, 1, inner_dim]
return residual + x * gate
return residual + x * gate
# adapted from Diffusers: https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/normalization.py
@@ -168,7 +159,7 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
raise NotImplementedError(f"Norm type {norm_type} not implemented")
def forward(self, residual: torch.Tensor, x: torch.Tensor,
gate: torch.Tensor | int, shift: torch.Tensor,
gate: torch.Tensor, shift: torch.Tensor,
scale: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""
Apply gated residual connection, followed by layernorm and
@@ -180,41 +171,12 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
- residual value (value after residual connection
but before normalization)
"""
# x.shape: [batch_size, seq_len, inner_dim]
# Apply residual connection with gating
if isinstance(gate, int):
# used by cross-attention, should be 1
assert gate == 1
residual_output = residual + x
elif isinstance(gate, torch.Tensor):
if gate.dim() == 4:
# gate.shape: [batch_size, num_frames, 1, inner_dim]
num_frames = gate.shape[1]
frame_seqlen = x.shape[1] // num_frames
residual_output = residual + (
x.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
gate).flatten(1, 2)
else:
# used by bidirectional self attention
# gate.shape: [batch_size, 1, inner_dim]
residual_output = residual + x * gate
else:
raise ValueError(f"Gate type {type(gate)} not supported")
# residual_output.shape: [batch_size, seq_len, inner_dim]
residual_output = residual + x * gate
# Apply normalization
normalized = self.norm(residual_output)
# Apply scale and shift
if isinstance(scale, torch.Tensor) and scale.dim() == 4:
# scale.shape: [batch_size, num_frames, 1, inner_dim]
# shift.shape: [batch_size, num_frames, 1, inner_dim]
num_frames = scale.shape[1]
frame_seqlen = normalized.shape[1] // num_frames
modulated = (
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1.0 + scale) + shift).flatten(1, 2)
else:
modulated = normalized * (1.0 + scale) + shift
modulated = normalized * (1.0 + scale) + shift
return modulated, residual_output
@@ -256,24 +218,8 @@ class LayerNormScaleShift(nn.Module):
def forward(self, x: torch.Tensor, shift: torch.Tensor,
scale: torch.Tensor) -> torch.Tensor:
"""Apply ln followed by scale and shift in a single fused operation."""
# x.shape: [batch_size, seq_len, inner_dim]
normalized = self.norm(x)
if self.compute_dtype == torch.float32:
normalized = normalized.float()
if scale.dim() == 4:
# scale.shape: [batch_size, num_frames, 1, inner_dim]
num_frames = scale.shape[1]
frame_seqlen = normalized.shape[1] // num_frames
output = (
normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1.0 + scale) + shift).flatten(1, 2)
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
else:
# scale.shape: [batch_size, 1, inner_dim]
# shift.shape: [batch_size, 1, inner_dim]
output = normalized * (1.0 + scale) + shift
if self.compute_dtype == torch.float32:
output = output.to(x.dtype)
return output
return normalized * (1.0 + scale) + shift
+1 -1
View File
@@ -63,7 +63,7 @@ class BaseLayerWithLoRA(nn.Module):
device=self.base_layer.weight.device,
dtype=self.base_layer.weight.dtype))
torch.nn.init.kaiming_uniform_(self.lora_A, a=math.sqrt(5))
torch.nn.init.zeros_(self.lora_B)
torch.nn.init.kaiming_uniform_(self.lora_B, a=math.sqrt(5))
else:
self.lora_A = None
self.lora_B = None
+11 -20
View File
@@ -244,26 +244,19 @@ class CausalWanTransformerBlock(nn.Module):
current_start: int = 0,
cache_start: int | None = None,
) -> torch.Tensor:
# hidden_states.shape: [batch_size, seq_length, inner_dim]
# temb.shape: [batch_size, num_frames, 6, inner_dim]
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
num_frames = temb.shape[1]
frame_seqlen = hidden_states.shape[1] // num_frames
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
# e.shape: [batch_size, num_frames, 6, inner_dim]
assert e.shape == (bs, num_frames, 6, self.hidden_dim)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=2)
# *_msa.shape: [batch_size, num_frames, 1, inner_dim]
6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).flatten(1, 2).to(orig_dtype)
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
@@ -494,8 +487,8 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = hidden_states.flatten(2).transpose(1, 2)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
timestep, encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
@@ -533,9 +526,8 @@ class CausalWanTransformer3DModel(BaseDiT):
**causal_kwargs)
# 5. Output norm, projection & unpatchify
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2,
dim=2)
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
@@ -604,8 +596,8 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = hidden_states.flatten(2).transpose(1, 2)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
timestep, encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
@@ -631,9 +623,8 @@ class CausalWanTransformer3DModel(BaseDiT):
block_mask=self.block_mask)
# 5. Output norm, projection & unpatchify
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2,
dim=2)
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
-3
View File
@@ -61,9 +61,6 @@ _SCHEDULERS = {
"FlowMatchEulerDiscreteScheduler"),
"UniPCMultistepScheduler":
("schedulers", "scheduling_unipc_multistep", "UniPCMultistepScheduler"),
"SelfForcingFlowMatchScheduler":
("schedulers", "scheduling_self_forcing_flow_match",
"SelfForcingFlowMatchScheduler"),
}
_FAST_VIDEO_MODELS = {
@@ -1,122 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from diffusers.configuration_utils import ConfigMixin, register_to_config
from diffusers.schedulers.scheduling_utils import SchedulerMixin
from diffusers.utils import BaseOutput
import torch
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.base import BaseScheduler
logger = init_logger(__name__)
class SelfForcingFlowMatchSchedulerOutput(BaseOutput):
"""
Output class for the scheduler's `step` function output.
Args:
prev_sample (`torch.FloatTensor` of shape `(batch_size, num_channels, height, width)` for images):
Computed sample `(x_{t-1})` of previous timestep. `prev_sample` should be used as next model input in the
denoising loop.
"""
prev_sample: torch.FloatTensor
class SelfForcingFlowMatchScheduler(BaseScheduler, ConfigMixin, SchedulerMixin):
order = 1
def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, inverse_timesteps=False, extra_one_step=False, reverse_sigmas=False, training=False):
self.num_train_timesteps = num_train_timesteps
self.shift = shift
self.sigma_max = sigma_max
self.sigma_min = sigma_min
self.inverse_timesteps = inverse_timesteps
self.extra_one_step = extra_one_step
self.reverse_sigmas = reverse_sigmas
self.set_timesteps(num_inference_steps, training=training)
def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, training=False, return_dict=False, **kwargs):
sigma_start = self.sigma_min + \
(self.sigma_max - self.sigma_min) * denoising_strength
if self.extra_one_step:
self.sigmas = torch.linspace(
sigma_start, self.sigma_min, num_inference_steps + 1)[:-1]
else:
self.sigmas = torch.linspace(
sigma_start, self.sigma_min, num_inference_steps)
if self.inverse_timesteps:
self.sigmas = torch.flip(self.sigmas, dims=[0])
self.sigmas = self.shift * self.sigmas / \
(1 + (self.shift - 1) * self.sigmas)
if self.reverse_sigmas:
self.sigmas = 1 - self.sigmas
self.timesteps = self.sigmas * self.num_train_timesteps
if training:
x = self.timesteps
y = torch.exp(-2 * ((x - num_inference_steps / 2) /
num_inference_steps) ** 2)
y_shifted = y - y.min()
bsmntw_weighing = y_shifted * \
(num_inference_steps / y_shifted.sum())
self.linear_timesteps_weights = bsmntw_weighing
def step(self, model_output: torch.FloatTensor, timestep: torch.FloatTensor, sample: torch.FloatTensor, to_final=False, return_dict=False, **kwargs):
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
self.sigmas = self.sigmas.to(model_output.device)
self.timesteps = self.timesteps.to(model_output.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
if to_final or (timestep_id + 1 >= len(self.timesteps)).any():
sigma_ = 1 if (
self.inverse_timesteps or self.reverse_sigmas) else 0
else:
sigma_ = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1)
prev_sample = sample + model_output * (sigma_ - sigma)
if isinstance(prev_sample, torch.Tensor | float) and not return_dict:
return (prev_sample, )
return SelfForcingFlowMatchSchedulerOutput(prev_sample=prev_sample)
def add_noise(self, original_samples, noise, timestep):
"""
Diffusion forward corruption process.
Input:
- clean_latent: the clean latent with shape [B*T, C, H, W]
- noise: the noise with shape [B*T, C, H, W]
- timestep: the timestep with shape [B*T]
Output: the corrupted latent with shape [B*T, C, H, W]
"""
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
self.sigmas = self.sigmas.to(noise.device)
self.timesteps = self.timesteps.to(noise.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
sample = (1 - sigma) * original_samples + sigma * noise
return sample.type_as(noise)
def training_target(self, sample, noise, timestep):
target = noise - sample
return target
def training_weight(self, timestep):
"""
Input:
- timestep: the timestep with shape [B*T]
Output: the corresponding weighting [B*T]
"""
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
self.linear_timesteps_weights = self.linear_timesteps_weights.to(timestep.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(1) - timestep.unsqueeze(0)).abs(), dim=0)
weights = self.linear_timesteps_weights[timestep_id]
return weights
def scale_model_input(self, sample: torch.Tensor, timestep: int | None = None) -> torch.Tensor:
return sample
def set_shift(self, shift: float) -> None:
self.shift = shift
+1 -23
View File
@@ -145,30 +145,8 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
scheduler: Any) -> torch.Tensor:
"""
Convert predicted noise to clean latent.
Args:
pred_noise: the predicted noise with shape [B, C, H, W]
where B is batch_size or batch_size * num_frames
noise_input_latent: the noisy latent with shape [B, C, H, W],
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
scheduler: the scheduler
Returns:
the predicted video with shape [B, C, H, W]
"""
# If timestep is [bs, num_frames]
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
assert timestep.numel() == noise_input_latent.shape[0]
elif timestep.ndim == 1:
# If timestep is [1]
if timestep.shape[0] == 1:
timestep = timestep.expand(noise_input_latent.shape[0])
else:
assert timestep.numel() == noise_input_latent.shape[0]
else:
raise ValueError(f"[pred_noise_to_pred_video] Invalid timestep shape: {timestep.shape}")
# timestep shape should be [B]
timestep = timestep.expand(noise_input_latent.shape[0])
dtype = pred_noise.dtype
device = pred_noise.device
pred_noise = pred_noise.float().to(device)
@@ -16,7 +16,8 @@ from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
CausalDMDDenosingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage)
TextEncodingStage,
TimestepPreparationStage)
# isort: on
logger = init_logger(__name__)
@@ -47,6 +48,10 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
@@ -40,7 +40,7 @@ class ComposedPipelineBase(ABC):
_extra_config_module_map: dict[str, str] = {}
training_args: TrainingArgs | None = None
fastvideo_args: FastVideoArgs | TrainingArgs | None = None
modules: dict[str, Any] = {}
modules: dict[str, torch.nn.Module] = {}
post_init_called: bool = False
# TODO(will): args should support both inference args and training args
@@ -79,7 +79,7 @@ class ComposedPipelineBase(ABC):
for name, module in self.modules.items():
if not isinstance(module, torch.nn.Module):
continue
if "transformer" in name:
if name == "transformer":
module.requires_grad_(True)
else:
module.requires_grad_(False)
+12 -36
View File
@@ -7,7 +7,7 @@ import torch
import torch.distributed as dist
import torch.nn as nn
from safetensors.torch import load_file
from torch.distributed.device_mesh import DeviceMesh, init_device_mesh
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.tensor import DTensor
from fastvideo.distributed import get_local_torch_device
@@ -32,7 +32,6 @@ class LoRAPipeline(ComposedPipelineBase):
cur_adapter_name: str = ""
cur_adapter_path: str = ""
lora_layers: dict[str, BaseLayerWithLoRA] = {}
lora_layers_critic: dict[str, BaseLayerWithLoRA] = {}
fastvideo_args: FastVideoArgs | TrainingArgs
exclude_lora_layers: list[str] = []
device: torch.device = get_local_torch_device()
@@ -82,17 +81,6 @@ class LoRAPipeline(ComposedPipelineBase):
def set_trainable(self) -> None:
def set_lora_grads(lora_layers: dict[str, BaseLayerWithLoRA],
device_mesh: DeviceMesh):
for name, layer in lora_layers.items():
layer.lora_A.requires_grad_(True)
layer.lora_B.requires_grad_(True)
layer.base_layer.requires_grad_(False)
layer.lora_A = nn.Parameter(
DTensor.from_local(layer.lora_A, device_mesh=device_mesh))
layer.lora_B = nn.Parameter(
DTensor.from_local(layer.lora_B, device_mesh=device_mesh))
is_lora_training = self.training_mode and getattr(
self.fastvideo_args, "lora_training", False)
if not is_lora_training:
@@ -100,12 +88,18 @@ class LoRAPipeline(ComposedPipelineBase):
return
self.modules["transformer"].requires_grad_(False)
if "fake_score_transformer" in self.modules:
self.modules["fake_score_transformer"].requires_grad_(False)
device_mesh = init_device_mesh("cuda", (dist.get_world_size(), 1),
mesh_dim_names=["fake", "replicate"])
set_lora_grads(self.lora_layers, device_mesh)
set_lora_grads(self.lora_layers_critic, device_mesh)
for name, layer in self.lora_layers.items():
# Enable grads for lora weights only
# Must convert to DTensor for compatibility with other FSDP modules in grad calculation
layer.lora_A.requires_grad_(True)
layer.lora_B.requires_grad_(True)
layer.base_layer.requires_grad_(False)
layer.lora_A = nn.Parameter(
DTensor.from_local(layer.lora_A, device_mesh=device_mesh))
layer.lora_B = nn.Parameter(
DTensor.from_local(layer.lora_B, device_mesh=device_mesh))
def convert_to_lora_layers(self) -> None:
"""
@@ -137,24 +131,6 @@ class LoRAPipeline(ComposedPipelineBase):
converted_count += 1
logger.info("Converted %d layers to LoRA layers", converted_count)
if "fake_score_transformer" in self.modules:
for name, layer in self.modules[
"fake_score_transformer"].named_modules():
if not self.is_target_layer(name):
continue
layer = get_lora_layer(layer,
lora_rank=self.lora_rank,
lora_alpha=self.lora_alpha,
training_mode=self.training_mode)
if layer is not None:
self.lora_layers_critic[name] = layer
replace_submodule(self.modules["fake_score_transformer"],
name, layer)
converted_count += 1
logger.info(
"Converted %d layers to LoRA layers in the critic model",
converted_count)
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None): # type: ignore
@@ -248,4 +224,4 @@ class LoRAPipeline(ComposedPipelineBase):
def unmerge_lora_weights(self) -> None:
for name, layer in self.lora_layers.items():
layer.unmerge_lora_weights()
layer.unmerge_lora_weights()
+1 -10
View File
@@ -147,12 +147,7 @@ class ForwardBatch:
modules: dict[str, Any] = field(default_factory=dict)
# Final output (after pipeline completion)
output: torch.Tensor | None = None
return_trajectory_latents: bool = False
return_trajectory_decoded: bool = False
trajectory_timesteps: list[torch.Tensor] | None = None
trajectory_latents: torch.Tensor | None = None
trajectory_decoded: list[torch.Tensor] | None = None
output: Any = None
# Extra parameters that might be needed by specific pipeline implementations
extra: dict[str, Any] = field(default_factory=dict)
@@ -211,10 +206,6 @@ class TrainingBatch:
infos: list[dict[str, Any]] | None = None
mask_lat_size: torch.Tensor | None = None
# ODE trajectory supervision
trajectory_latents: torch.Tensor | None = None
trajectory_timesteps: torch.Tensor | None = None
# Transformer inputs
noisy_model_input: torch.Tensor | None = None
timesteps: torch.Tensor | None = None
@@ -1,5 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import multiprocessing
import os
from concurrent.futures import ProcessPoolExecutor
from typing import Any
import numpy as np
@@ -17,8 +19,6 @@ from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages import TextEncodingStage
from fastvideo.workflow.preprocess.parquet_io import (ParquetDatasetWriter,
records_to_table)
logger = init_logger(__name__)
@@ -54,13 +54,9 @@ class BasePreprocessPipeline(ComposedPipelineBase):
"""Get additional features specific to the pipeline type. Override in subclasses."""
return {}
def get_pyarrow_schema(self) -> pa.Schema:
"""Return the PyArrow schema for this pipeline. Must be overridden."""
raise NotImplementedError
def get_schema_fields(self) -> list[str]:
"""Get the schema fields for the pipeline type."""
return [f.name for f in self.get_pyarrow_schema()]
"""Get the schema fields for the pipeline type. Override in subclasses."""
raise NotImplementedError
def create_record_for_schema(self,
preprocess_batch: PreprocessBatch,
@@ -404,26 +400,166 @@ class BasePreprocessPipeline(ComposedPipelineBase):
batch_data.append(record)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
table = records_to_table(batch_data, self.get_pyarrow_schema())
# Convert batch data to PyArrow arrays
arrays = []
for field in self.get_schema_fields():
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
elif field in ['width', 'height', 'num_frames']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.int32()))
elif field in ['duration_sec', 'fps']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.float32()))
else:
arrays.append(
pa.array([record[field] for record in batch_data]))
table = pa.Table.from_arrays(arrays,
names=self.get_schema_fields())
write_pbar.update(1)
write_pbar.close()
if not hasattr(self, 'dataset_writer'):
self.dataset_writer = ParquetDatasetWriter(
out_dir=combined_parquet_dir,
samples_per_file=args.samples_per_file,
)
self.dataset_writer.append_table(table)
# Store the table in a list for later processing
if not hasattr(self, 'all_tables'):
self.all_tables = []
self.all_tables.append(table)
logger.info("Collected batch with %s samples", len(table))
if num_processed_samples >= args.flush_frequency:
written = self.dataset_writer.flush()
logger.info("Flushed %s samples to parquet", written)
self._flush_tables(num_processed_samples, args,
combined_parquet_dir)
num_processed_samples = 0
self.all_tables = []
def _final_flush_if_any(self):
if hasattr(self, 'dataset_writer'):
self.dataset_writer.flush()
def _flush_tables(self, num_processed_samples: int, args,
combined_parquet_dir: str):
"""Flush collected tables to disk."""
assert hasattr(self, 'all_tables') and self.all_tables
print(f"Combining {len(self.all_tables)} batches...")
combined_table = pa.concat_tables(self.all_tables)
assert len(combined_table) == num_processed_samples
print(f"Total samples collected: {len(combined_table)}")
# Calculate total number of chunks needed, discarding remainder
total_chunks = max(num_processed_samples // args.samples_per_file, 1)
print(f"Fixed samples per parquet file: {args.samples_per_file}")
print(f"Total number of parquet files: {total_chunks}")
print(
f"Total samples to be processed: {total_chunks * args.samples_per_file} (discarding {num_processed_samples % args.samples_per_file} samples)"
)
# Split work among processes
num_workers = int(min(multiprocessing.cpu_count(), total_chunks))
chunks_per_worker = (total_chunks + num_workers - 1) // num_workers
print(f"Using {num_workers} workers to process {total_chunks} chunks")
logger.info("Chunks per worker: %s", chunks_per_worker)
# Prepare work ranges
work_ranges = []
for i in range(num_workers):
start_idx = i * chunks_per_worker
end_idx = min((i + 1) * chunks_per_worker, total_chunks)
if start_idx < total_chunks:
work_ranges.append(
(start_idx, end_idx, combined_table, i,
combined_parquet_dir, args.samples_per_file))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=num_workers) as executor:
futures = {
executor.submit(self.process_chunk_range, work_range):
work_range
for work_range in work_ranges
}
for future in tqdm(futures, desc="Processing chunks"):
try:
written = future.result()
total_written += written
logger.info("Processed chunk with %s samples", written)
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error("Failed to process range %s-%s: %s",
work_range[0], work_range[1], str(e))
# Retry failed ranges sequentially
if failed_ranges:
logger.warning("Retrying %s failed ranges sequentially",
len(failed_ranges))
for work_range in failed_ranges:
try:
total_written += self.process_chunk_range(work_range)
except Exception as e:
logger.error(
"Failed to process range %s-%s after retry: %s",
work_range[0], work_range[1], str(e))
logger.info("Total samples written: %s", total_written)
@staticmethod
def process_chunk_range(args: Any) -> int:
start_idx, end_idx, table, worker_id, output_dir, samples_per_file = args
try:
total_written = 0
num_samples = len(table)
# Create worker-specific subdirectory
worker_dir = os.path.join(output_dir, f"worker_{worker_id}")
os.makedirs(worker_dir, exist_ok=True)
# Check how many files there are already in the dir, and update i accordingly
num_parquets = 0
for root, _, files in os.walk(worker_dir):
for file in files:
if file.endswith('.parquet'):
num_parquets += 1
for i in range(start_idx, end_idx):
start_sample = i * samples_per_file
end_sample = min((i + 1) * samples_per_file, num_samples)
chunk = table.slice(start_sample, end_sample - start_sample)
# Create chunk file in worker's directory
chunk_path = os.path.join(
worker_dir, f"data_chunk_{i + num_parquets}.parquet")
temp_path = chunk_path + '.tmp'
try:
# Write to temporary file
pq.write_table(chunk, temp_path, compression='zstd')
# Rename temporary file to final file
if os.path.exists(chunk_path):
os.remove(
chunk_path) # Remove existing file if it exists
os.rename(temp_path, chunk_path)
total_written += len(chunk)
except Exception as e:
# Clean up temporary file if it exists
if os.path.exists(temp_path):
os.remove(temp_path)
raise e
return total_written
except Exception as e:
logger.error("Error processing chunks %s-%s for worker %s: %s",
start_idx, end_idx, worker_id, str(e))
raise
@@ -40,9 +40,9 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
image_processor=self.get_module("image_processor"),
))
def get_pyarrow_schema(self):
"""Return the PyArrow schema for I2V pipeline."""
return pyarrow_schema_i2v
def get_schema_fields(self) -> list[str]:
"""Get the schema fields for I2V pipeline."""
return [f.name for f in pyarrow_schema_i2v]
def get_extra_features(self, valid_data: dict[str, Any],
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
@@ -15,29 +15,27 @@ from typing import Any
import numpy as np
import pyarrow as pa
import torch
from PIL import Image
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm import tqdm
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import gettextdataset
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_ode_trajectory_text_only)
from fastvideo.dataset import getdataset, gettextdataset
from fastvideo.dataset.dataloader.schema import pyarrow_schema_ode_trajectory, pyarrow_schema_ode_trajectory_text_only
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
SelfForcingFlowMatchScheduler)
from fastvideo.utils import shallow_asdict, save_decoded_latents_as_video
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
from fastvideo.pipelines.stages import (DenoisingStage, ImageVAEEncodingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
from fastvideo.utils import save_decoded_latents_as_video, shallow_asdict
from fastvideo.workflow.preprocess.parquet_io import (ParquetDatasetWriter,
records_to_table)
TimestepPreparationStage,
DecodingStage)
logger = init_logger(__name__)
@@ -50,36 +48,16 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
]
preprocess_dataloader: StatefulDataLoader
preprocess_loader_iter: Iterator[dict[str, Any]]
pbar: Any
num_processed_samples: int
def get_pyarrow_schema(self) -> pa.Schema:
"""Return the PyArrow schema for ODE Trajectory pipeline."""
return pyarrow_schema_ode_trajectory_text_only
def get_schema_fields(self):
"""Get the schema fields for ODE Trajectory pipeline."""
# Check if we're using text dataset by checking if the dataset is TextDataset
if hasattr(self, 'preprocess_dataloader') and hasattr(self.preprocess_dataloader.dataset, '_process_text_data'):
return [f.name for f in pyarrow_schema_ode_trajectory_text_only]
return [f.name for f in pyarrow_schema_ode_trajectory]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
logger.info('WTF flow_shift: %s',
fastvideo_args.pipeline_config.flow_shift)
assert fastvideo_args.pipeline_config.flow_shift == 5
# self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
# shift=fastvideo_args.pipeline_config.flow_shift)
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
shift=fastvideo_args.pipeline_config.flow_shift,
sigma_min=0.0,
extra_one_step=True)
self.modules["scheduler"].set_timesteps(num_inference_steps=48,
denoising_strength=1.0)
# logger.info('WTF scheduler timesteps: %s',
# self.modules["scheduler"].timesteps)
# scheduler = FlowMatchScheduler(
# shift=8.0, sigma_min=0.0, extra_one_step=True)
# device = get_local_torch_device()
# # scheduler.num_train_timesteps = 100
# scheduler.set_timesteps(num_inference_steps=50, denoising_strength=1.0)
# scheduler.sigmas = scheduler.sigmas.to(device)
# self.modules["scheduler"] = scheduler
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
@@ -87,6 +65,9 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="vae_encoding_stage",
stage=ImageVAEEncodingStage(
vae=self.get_module("vae"), ))
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
@@ -97,16 +78,272 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
pipeline=self,
))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def preprocess_text_and_trajectory(self, fastvideo_args: FastVideoArgs,
args):
"""Preprocess text-only data and generate trajectory information."""
def preprocess_video_and_text_and_trajectory(self,
fastvideo_args: FastVideoArgs,
args):
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
# Filter out invalid samples (those with all zeros)
valid_indices = []
for i, pixel_values in enumerate(data["pixel_values"]):
if not torch.all(
pixel_values == 0): # Check if all values are zero
valid_indices.append(i)
self.num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples
valid_data = {
"pixel_values":
torch.stack(
[data["pixel_values"][i] for i in valid_indices]),
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
"fps": [data["fps"][i] for i in valid_indices],
"duration": [data["duration"][i] for i in valid_indices],
}
# VAE
with torch.autocast("cuda", dtype=torch.float32):
latents = self.get_module("vae").encode(
valid_data["pixel_values"].to(
get_local_torch_device())).mean
# Get extra features if needed
extra_features = self.get_extra_features(
valid_data, fastvideo_args)
batch_captions = valid_data["text"]
# Encode text using the standalone TextEncodingStage API
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
# # Get sequence lengths from attention masks (number of 1s)
# seq_lens = prompt_attention_mask.sum(dim=1)
# non_padded_embeds = []
# non_padded_masks = []
# # Process each item in the batch
# for i in range(prompt_embeds.size(0)):
# seq_len = seq_lens[i].item()
# # Slice the embeddings and masks to keep only non-padding parts
# non_padded_embeds.append(prompt_embeds[i, :seq_len])
# non_padded_masks.append(prompt_attention_mask[i, :seq_len])
# Update the tensors with non-padded versions
# prompt_embeds = non_padded_embeds
# prompt_attention_masks = non_padded_masks
# prompt_embeds = prompt_embeds
# logger.info(f"===== prompt_embeds: {prompt_embeds[0].shape}")
# logger.info(f"===== prompt_attention_masks: {prompt_attention_masks[0].shape}")
sampling_params = SamplingParam.from_pretrained(
args.model_path)
# encode negative prompt for trajectory collection
if sampling_params.guidance_scale > 1 and sampling_params.negative_prompt is not None:
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
sampling_params.negative_prompt,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
negative_prompt_embed = negative_prompt_embeds_list[0][0]
negative_prompt_attention_mask = negative_prompt_masks_list[0][0]
else:
negative_prompt_embed = None
negative_prompt_attention_mask = None
trajectory_latents = []
trajectory_timesteps = []
trajectory_decoded = []
for i, (prompt_embed, prompt_attention_mask) in enumerate(zip(prompt_embeds, prompt_attention_masks)):
prompt_embed = prompt_embed.unsqueeze(0)
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
logger.info(f"what")
logger.info(f"===== prompt_embed: {prompt_embed.shape}")
logger.info(f"===== prompt_attention_mask: {prompt_attention_mask.shape}")
# Collect the trajectory data
batch = ForwardBatch(
**shallow_asdict(sampling_params),
# data_type="video",
# seed=args.seed,
# prompt=batch_captions[i],
# prompt_embeds=[prompt_embed],
# prompt_attention_mask=[prompt_attention_mask],
# height=args.max_height,
# width=args.max_width,
# num_frames=81,
# fps=args.train_fps,
# return_trajectory_latents=True,
# guidance_scale=3.0,
# do_classifier_free_guidance=True,
)
batch.prompt_embeds = [prompt_embed]
batch.prompt_attention_mask = [prompt_attention_mask]
batch.negative_prompt_embeds = [negative_prompt_embed]
batch.negative_attention_mask = [negative_prompt_attention_mask]
batch.return_trajectory_latents = True
batch.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
# batch.num_frames = 81
batch.fps = args.train_fps
batch.guidance_scale = 3.0
batch.do_classifier_free_guidance = True
# fastvideo_args.pipeline_config.ti2v_task = True
result_batch = self.input_validation_stage(
batch, fastvideo_args)
# result_batch = self.prompt_encoding_stage(result_batch, fastvideo_args)
# result_batch = self.vae_encoding_stage(result_batch, fastvideo_args)
result_batch = self.timestep_preparation_stage(
batch, fastvideo_args)
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch, fastvideo_args)
# trajectory_latents = result_batch.trajectory_latents
trajectory_latents.append(result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(result_batch.trajectory_timesteps.cpu())
trajectory_decoded.append(result_batch.trajectory_decoded)
extra_features["trajectory_latents"] = trajectory_latents
extra_features["trajectory_timesteps"] = trajectory_timesteps
logger.info(f"===== trajectory_latents: {trajectory_latents[0].shape}")
logger.info(f"===== trajectory_latents len: {len(trajectory_latents)}")
logger.info(f"===== trajectory_timesteps: {trajectory_timesteps}")
logger.info(f"===== trajectory_timesteps len: {len(trajectory_timesteps)}")
if batch.return_trajectory_decoded:
logger.info(f"===== SAVING TRAJECTORY DECODED")
for i, decoded_frames in enumerate(trajectory_decoded):
for j, decoded_frame in enumerate(decoded_frames):
logger.info(f"===== SAVING TRAJECTORY DECODED {i} for prompt {batch_captions[i]}")
save_decoded_latents_as_video(decoded_frame, f"decoded_videos/trajectory_decoded_{i}_{j}.mp4", args.train_fps)
# assert False
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, video_path in save_pbar:
# Get the corresponding latent and info using video name
latent = latents[idx].cpu()
video_name = os.path.basename(video_path).split(".")[0]
# Convert tensors to numpy arrays
vae_latent = latent.cpu().numpy()
text_embedding = prompt_embeds[idx].cpu().numpy()
# Get extra features for this sample if needed
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
logger.info(f"===== key: {key}")
if isinstance(value, torch.Tensor):
logger.info(f"===== value: {value[idx].shape}")
sample_extra_features[key] = value[idx].cpu().numpy(
)
else:
assert isinstance(value, list)
if isinstance(value[idx], torch.Tensor):
logger.info(f"===== value in list: {value[idx].shape}")
sample_extra_features[key] = value[idx].cpu().float().numpy(
)
else:
logger.info(f"===== value in list: not tensor")
sample_extra_features[key] = value[idx]
# logger.info(f"===== value: not tensor")
# sample_extra_features[key] = value[idx]
# Create record for Parquet dataset
record = self.create_record(
video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=sample_extra_features)
batch_data.append(record)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = []
for field in self.get_schema_fields():
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
elif field in ['width', 'height', 'num_frames']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.int32()))
elif field in ['duration_sec', 'fps']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.float32()))
else:
arrays.append(
pa.array([record[field] for record in batch_data]))
table = pa.Table.from_arrays(arrays,
names=self.get_schema_fields())
write_pbar.update(1)
write_pbar.close()
# Store the table in a list for later processing
if not hasattr(self, 'all_tables'):
self.all_tables = []
self.all_tables.append(table)
logger.info("Collected batch with %s samples", len(table))
if self.num_processed_samples >= args.flush_frequency:
self._flush_tables(self.num_processed_samples, args,
self.combined_parquet_dir)
self.num_processed_samples = 0
self.all_tables = []
def preprocess_text_and_trajectory(self,
fastvideo_args: FastVideoArgs,
args):
"""Preprocess text-only data and generate trajectory information."""
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
@@ -128,14 +365,12 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
}
# Add fps and duration if available in data
if "fps" in data:
valid_data["fps"] = [data["fps"][i] for i in valid_indices]
if "duration" in data:
valid_data["duration"] = [
data["duration"][i] for i in valid_indices
]
valid_data["duration"] = [data["duration"][i] for i in valid_indices]
batch_captions = valid_data["text"]
# Encode text using the standalone TextEncodingStage API
@@ -149,7 +384,8 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
sampling_params = SamplingParam.from_pretrained(args.model_path)
sampling_params = SamplingParam.from_pretrained(
args.model_path)
# encode negative prompt for trajectory collection
if sampling_params.guidance_scale > 1 and sampling_params.negative_prompt is not None:
@@ -160,8 +396,7 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
return_attention_mask=True,
)
negative_prompt_embed = negative_prompt_embeds_list[0][0]
negative_prompt_attention_mask = negative_prompt_masks_list[
0][0]
negative_prompt_attention_mask = negative_prompt_masks_list[0][0]
else:
negative_prompt_embed = None
negative_prompt_attention_mask = None
@@ -169,28 +404,25 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
trajectory_latents = []
trajectory_timesteps = []
trajectory_decoded = []
for i, (prompt_embed, prompt_attention_mask) in enumerate(
zip(prompt_embeds, prompt_attention_masks,
strict=False)):
for i, (prompt_embed, prompt_attention_mask) in enumerate(zip(prompt_embeds, prompt_attention_masks)):
prompt_embed = prompt_embed.unsqueeze(0)
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
# Collect the trajectory data (text-to-video generation)
batch = ForwardBatch(**shallow_asdict(sampling_params), )
batch = ForwardBatch(
**shallow_asdict(sampling_params),
)
batch.prompt_embeds = [prompt_embed]
batch.prompt_attention_mask = [prompt_attention_mask]
batch.negative_prompt_embeds = [negative_prompt_embed]
batch.negative_attention_mask = [
negative_prompt_attention_mask
]
batch.num_inference_steps = 48
batch.negative_attention_mask = [negative_prompt_attention_mask]
batch.return_trajectory_latents = True
batch.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
batch.fps = args.train_fps
batch.guidance_scale = 6.0
batch.guidance_scale = 3.0
batch.do_classifier_free_guidance = True
result_batch = self.input_validation_stage(
@@ -201,13 +433,10 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch,
fastvideo_args)
trajectory_latents.append(
result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(
result_batch.trajectory_timesteps.cpu())
result_batch = self.decoding_stage(result_batch, fastvideo_args)
trajectory_latents.append(result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(result_batch.trajectory_timesteps.cpu())
trajectory_decoded.append(result_batch.trajectory_decoded)
# Prepare extra features for text-only processing
@@ -216,13 +445,17 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"trajectory_timesteps": trajectory_timesteps
}
logger.info(f"===== trajectory_latents: {trajectory_latents[0].shape}")
logger.info(f"===== trajectory_latents len: {len(trajectory_latents)}")
logger.info(f"===== trajectory_timesteps: {trajectory_timesteps}")
logger.info(f"===== trajectory_timesteps len: {len(trajectory_timesteps)}")
if batch.return_trajectory_decoded:
logger.info(f"===== SAVING TRAJECTORY DECODED")
for i, decoded_frames in enumerate(trajectory_decoded):
for j, decoded_frame in enumerate(decoded_frames):
save_decoded_latents_as_video(
decoded_frame,
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
args.train_fps)
logger.info(f"===== SAVING TRAJECTORY DECODED {i} for prompt {batch_captions[i]}")
save_decoded_latents_as_video(decoded_frame, f"decoded_videos/trajectory_decoded_{i}_{j}.mp4", args.train_fps)
# Prepare batch data for Parquet dataset
batch_data = []
@@ -232,7 +465,7 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
desc="Saving outputs",
unit="item",
leave=False)
for idx, video_path in save_pbar:
video_name = os.path.basename(video_path).split(".")[0]
@@ -243,15 +476,17 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
logger.info(f"===== key: {key}")
if isinstance(value, torch.Tensor):
sample_extra_features[key] = value[idx].cpu(
).numpy()
logger.info(f"===== value: {value[idx].shape}")
sample_extra_features[key] = value[idx].cpu().numpy()
else:
assert isinstance(value, list)
if isinstance(value[idx], torch.Tensor):
sample_extra_features[key] = value[idx].cpu(
).float().numpy()
logger.info(f"===== value in list: {value[idx].shape}")
sample_extra_features[key] = value[idx].cpu().float().numpy()
else:
logger.info(f"===== value in list: not tensor")
sample_extra_features[key] = value[idx]
# Create record for Parquet dataset (without VAE latents for text-only)
@@ -265,33 +500,57 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
batch_data.append(record)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
table = records_to_table(batch_data,
self.get_pyarrow_schema())
# Convert batch data to PyArrow arrays
arrays = []
for field in self.get_schema_fields():
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
elif field in ['width', 'height', 'num_frames']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.int32()))
elif field in ['duration_sec', 'fps']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.float32()))
else:
arrays.append(
pa.array([record[field] for record in batch_data]))
table = pa.Table.from_arrays(arrays,
names=self.get_schema_fields())
write_pbar.update(1)
write_pbar.close()
if not hasattr(self, 'dataset_writer'):
self.dataset_writer = ParquetDatasetWriter(
out_dir=self.combined_parquet_dir,
samples_per_file=args.samples_per_file,
)
self.dataset_writer.append_table(table)
# Store the table in a list for later processing
if not hasattr(self, 'all_tables'):
self.all_tables = []
self.all_tables.append(table)
logger.info("Collected batch with %s samples", len(table))
if self.num_processed_samples >= args.flush_frequency:
written = self.dataset_writer.flush()
logger.info("Flushed %s samples to parquet", written)
self._flush_tables(self.num_processed_samples, args,
self.combined_parquet_dir)
self.num_processed_samples = 0
self.all_tables = []
# Final flush for any remaining samples
if hasattr(self, 'dataset_writer'):
written = self.dataset_writer.flush()
if written:
logger.info("Final flush wrote %s samples", written)
if hasattr(self, 'all_tables') and self.all_tables and self.num_processed_samples > 0:
logger.info(f"Final flush with {self.num_processed_samples} remaining samples")
self._flush_tables(self.num_processed_samples, args, self.combined_parquet_dir)
self.num_processed_samples = 0
self.all_tables = []
def create_text_only_record(
self,
@@ -302,7 +561,7 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
idx: int,
extra_features: dict[str, Any] | None = None) -> dict[str, Any]:
"""Create a record for text-only preprocessing using text-only schema."""
# Create base record using only fields from text-only schema
record = {
"id": f"text_{video_name}_{idx}",
@@ -314,23 +573,13 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"media_type": "text",
}
assert extra_features is not None, "extra_features is required"
assert "trajectory_latents" in extra_features, "trajectory_latents is required"
assert "trajectory_timesteps" in extra_features, "trajectory_timesteps is required"
# Add trajectory data if available
if extra_features and "trajectory_latents" in extra_features:
trajectory_latents = extra_features[
"trajectory_latents"][idx] if isinstance(
extra_features["trajectory_latents"],
list) else extra_features["trajectory_latents"]
trajectory_latents = extra_features["trajectory_latents"][idx] if isinstance(extra_features["trajectory_latents"], list) else extra_features["trajectory_latents"]
record.update({
"trajectory_latents_bytes":
trajectory_latents.tobytes(),
"trajectory_latents_shape":
list(trajectory_latents.shape),
"trajectory_latents_dtype":
str(trajectory_latents.dtype),
"trajectory_latents_bytes": trajectory_latents.tobytes(),
"trajectory_latents_shape": list(trajectory_latents.shape),
"trajectory_latents_dtype": str(trajectory_latents.dtype),
})
else:
record.update({
@@ -340,17 +589,11 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
})
if extra_features and "trajectory_timesteps" in extra_features:
trajectory_timesteps = extra_features[
"trajectory_timesteps"][idx] if isinstance(
extra_features["trajectory_timesteps"],
list) else extra_features["trajectory_timesteps"]
trajectory_timesteps = extra_features["trajectory_timesteps"][idx] if isinstance(extra_features["trajectory_timesteps"], list) else extra_features["trajectory_timesteps"]
record.update({
"trajectory_timesteps_bytes":
trajectory_timesteps.tobytes(),
"trajectory_timesteps_shape":
list(trajectory_timesteps.shape),
"trajectory_timesteps_dtype":
str(trajectory_timesteps.dtype),
"trajectory_timesteps_bytes": trajectory_timesteps.tobytes(),
"trajectory_timesteps_shape": list(trajectory_timesteps.shape),
"trajectory_timesteps_dtype": str(trajectory_timesteps.dtype),
})
else:
record.update({
@@ -361,6 +604,124 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
return record
def get_extra_features(self, valid_data: dict[str, Any],
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
# TODO(will): move these to cpu at some point
self.get_module("vae").to(get_local_torch_device())
# generator = torch.Generator(device=get_local_torch_device(), seed=42)
generator = torch.Generator("cpu").manual_seed(42)
features = {}
"""Get CLIP features from the first frame of each video."""
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
_, _, num_frames, height, width = valid_data["pixel_values"].shape
# latent_height = height // self.get_module(
# "vae").spatial_compression_ratio
# latent_width = width // self.get_module("vae").spatial_compression_ratio
unprocessed_images = []
pil_images = []
# Frame has values between -1 and 1
for frame in first_frame:
frame = (frame + 1) * 127.5
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
pil_images.append(frame_pil)
# processed_img = self.get_module("image_processor")(
# images=frame_pil, return_tensors="pt")
unprocessed_images.append(frame_pil)
"""Get VAE features from the first frame of each video"""
video_conditions = []
for frame in unprocessed_images:
latent = self.vae_encoding_stage.encode_image(
frame, height, width, fastvideo_args, generator)
video_conditions.append(latent)
features["image_condition_latents"] = video_conditions
features["pil_images"] = pil_images
return features
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
valid_data: dict[str, Any],
idx: int,
extra_features: dict[str, Any] | None = None) -> dict[str, Any]:
"""Create a record for the Parquet dataset with CLIP features."""
record = super().create_record(video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
if extra_features and "image_condition_latents" in extra_features:
image_condition_latents = extra_features["image_condition_latents"]
record.update({
"image_condition_latents_bytes":
image_condition_latents.tobytes(),
"image_condition_latents_shape":
list(image_condition_latents.shape),
"image_condition_latents_dtype":
str(image_condition_latents.dtype),
})
else:
record.update({
"image_condition_latents_bytes": b"",
"image_condition_latents_shape": [],
"image_condition_latents_dtype": "",
})
if extra_features and "trajectory_latents" in extra_features:
trajectory_latents = extra_features["trajectory_latents"]
record.update({
"trajectory_latents_bytes": trajectory_latents.tobytes(),
"trajectory_latents_shape": list(trajectory_latents.shape),
"trajectory_latents_dtype": str(trajectory_latents.dtype),
})
else:
record.update({
"trajectory_latents_bytes": b"",
"trajectory_latents_shape": [],
"trajectory_latents_dtype": "",
})
if extra_features and "trajectory_timesteps" in extra_features:
trajectory_timesteps = extra_features["trajectory_timesteps"]
record.update({
"trajectory_timesteps_bytes": trajectory_timesteps.tobytes(),
"trajectory_timesteps_shape": list(trajectory_timesteps.shape),
"trajectory_timesteps_dtype": str(trajectory_timesteps.dtype),
})
else:
record.update({
"trajectory_timesteps_bytes": b"",
"trajectory_timesteps_shape": [],
"trajectory_timesteps_dtype": "",
})
if extra_features and "pil_image" in extra_features:
pil_image = extra_features["pil_image"]
record.update({
"pil_image_bytes": pil_image.tobytes(),
"pil_image_shape": list(pil_image.shape),
"pil_image_dtype": str(pil_image.dtype),
})
else:
record.update({
"pil_image_bytes": b"",
"pil_image_shape": [],
"pil_image_dtype": "",
})
return record
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
if not self.post_init_called:
self.post_init()
@@ -373,6 +734,7 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
os.makedirs(self.combined_parquet_dir, exist_ok=True)
# Loading dataset
#train_dataset = getdataset(args)
train_dataset = gettextdataset(args)
self.preprocess_dataloader = DataLoader(
@@ -393,7 +755,8 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
# Initialize class variables for data sharing
self.video_data: dict[str, Any] = {} # Store video metadata and paths
self.latent_data: dict[str, Any] = {} # Store latent tensors
#self.preprocess_video_and_text_and_trajectory(fastvideo_args, args)
self.preprocess_text_and_trajectory(fastvideo_args, args)
EntryClass = PreprocessPipeline_ODE_Trajectory
EntryClass = PreprocessPipeline_ODE_Trajectory
@@ -15,9 +15,9 @@ class PreprocessPipeline_T2V(BasePreprocessPipeline):
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
def get_pyarrow_schema(self):
"""Return the PyArrow schema for T2V pipeline."""
return pyarrow_schema_t2v
def get_schema_fields(self):
"""Get the schema fields for T2V pipeline."""
return [f.name for f in pyarrow_schema_t2v]
EntryClass = PreprocessPipeline_T2V
+34 -12
View File
@@ -13,6 +13,7 @@ from fastvideo.pipelines.preprocess.preprocess_pipeline_ode_trajectory import (
PreprocessPipeline_ODE_Trajectory)
from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
PreprocessPipeline_T2V)
from fastvideo.pipelines.preprocess_text import PreprocessPipeline_Text
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
@@ -23,13 +24,22 @@ def main(args) -> None:
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 = {
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
"flow_shift": 5,
}
pipeline_config.update_config_from_dict(kwargs)
if args.preprocess_task == "text_only":
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
"text_encoder_cpu_offload": False,
}
pipeline_config.update_config_from_dict(kwargs)
else:
# Full config for video/image processing
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
}
pipeline_config.update_config_from_dict(kwargs)
fastvideo_args = FastVideoArgs(
model_path=args.model_path,
num_gpus=get_world_size(),
@@ -38,17 +48,25 @@ def main(args) -> None:
text_encoder_cpu_offload=False,
pipeline_config=pipeline_config,
)
if args.preprocess_task == "t2v":
PreprocessPipeline = PreprocessPipeline_T2V
elif args.preprocess_task == "i2v":
PreprocessPipeline = PreprocessPipeline_I2V
elif args.preprocess_task == "ode_trajectory":
print("Preprocess pipeline...")
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
elif args.preprocess_task == "text_only":
print("Text-only preprocessing pipeline...")
PreprocessPipeline = PreprocessPipeline_Text
else:
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}")
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}. "
f"Valid options: t2v, i2v, ode_trajectory, text_only")
logger.info(
f"Preprocess task: {args.preprocess_task} using {PreprocessPipeline.__name__}"
)
logger.info("Preprocess task: %s using %s", args.preprocess_task,
PreprocessPipeline.__name__)
pipeline = PreprocessPipeline(args.model_path, fastvideo_args)
pipeline.forward(batch=None, fastvideo_args=fastvideo_args, args=args)
@@ -87,7 +105,11 @@ if __name__ == "__main__":
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--preprocess_task", type=str, default="t2v")
parser.add_argument("--preprocess_task",
type=str,
default="t2v",
choices=["t2v", "i2v", "ode_trajectory", "text_only"],
help="Type of preprocessing task to run")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
parser.add_argument("--text_max_length", type=int, default=256)
@@ -111,4 +133,4 @@ if __name__ == "__main__":
)
args = parser.parse_args()
main(args)
main(args)
+216
View File
@@ -0,0 +1,216 @@
# SPDX-License-Identifier: Apache-2.0
"""
Text-only Data Preprocessing pipeline implementation.
This module contains an implementation of the Text-only Data Preprocessing pipeline
using the modular pipeline architecture, based on the ODE Trajectory preprocessing.
"""
import os
from collections.abc import Iterator
from typing import Any
import numpy as np
import pyarrow as pa
import torch
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm import tqdm
from fastvideo.dataset import gettextdataset
from fastvideo.dataset.dataloader.schema import pyarrow_schema_text_only
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import (TextEncodingStage)
logger = init_logger(__name__)
class PreprocessPipeline_Text(BasePreprocessPipeline):
"""Text-only preprocessing pipeline implementation."""
_required_config_modules = [
"text_encoder", "tokenizer"
]
preprocess_dataloader: StatefulDataLoader
preprocess_loader_iter: Iterator[dict[str, Any]]
def get_schema_fields(self):
"""Get the schema fields for text-only pipeline."""
return [f.name for f in pyarrow_schema_text_only]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
def preprocess_text_only(self,
fastvideo_args: FastVideoArgs,
args):
"""Preprocess text-only data."""
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
for i, text in enumerate(data["text"]):
if text and text.strip(): # Check if text is not empty
valid_indices.append(i)
self.num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples (text-only)
valid_data = {
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
}
batch_captions = valid_data["text"]
# Encode text using the standalone TextEncodingStage API
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
logger.info(f"===== prompt_embeds: {prompt_embeds.shape}")
logger.info(f"===== prompt_attention_masks: {prompt_attention_masks.shape}")
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, text_path in save_pbar:
text_name = os.path.basename(text_path).split(".")[0]
# Convert tensors to numpy arrays
text_embedding = prompt_embeds[idx].cpu().numpy()
# Create record for Parquet dataset (text-only)
record = self.create_text_only_record(
text_name=text_name,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx)
batch_data.append(record)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = []
for field in self.get_schema_fields():
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
else:
arrays.append(
pa.array([record[field] for record in batch_data]))
table = pa.Table.from_arrays(arrays,
names=self.get_schema_fields())
write_pbar.update(1)
write_pbar.close()
# Store the table in a list for later processing
if not hasattr(self, 'all_tables'):
self.all_tables = []
self.all_tables.append(table)
logger.info("Collected batch with %s samples", len(table))
if self.num_processed_samples >= args.flush_frequency:
self._flush_tables(self.num_processed_samples, args,
self.combined_parquet_dir)
self.num_processed_samples = 0
self.all_tables = []
# Final flush for any remaining samples
if hasattr(self, 'all_tables') and self.all_tables and self.num_processed_samples > 0:
logger.info(f"Final flush with {self.num_processed_samples} remaining samples")
self._flush_tables(self.num_processed_samples, args, self.combined_parquet_dir)
self.num_processed_samples = 0
self.all_tables = []
def create_text_only_record(
self,
text_name: str,
text_embedding: np.ndarray,
valid_data: dict[str, Any],
idx: int) -> dict[str, Any]:
"""Create a record for text-only preprocessing using text-only schema."""
# Create base record using only fields from text-only schema
record = {
"id": f"text_{text_name}_{idx}",
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
}
return record
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
if not self.post_init_called:
self.post_init()
self.local_rank = int(os.getenv("RANK", 0))
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
self.combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(self.combined_parquet_dir, exist_ok=True)
# Loading text dataset
train_dataset = gettextdataset(args)
self.preprocess_dataloader = DataLoader(
train_dataset,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
self.num_processed_samples = 0
# Add progress bar for text preprocessing
self.pbar = tqdm(self.preprocess_loader_iter,
desc="Processing text",
unit="batch",
disable=self.local_rank != 0)
# Initialize class variables for data sharing
self.text_data: dict[str, Any] = {} # Store text metadata and paths
self.preprocess_text_only(fastvideo_args, args)
EntryClass = PreprocessPipeline_Text
+7 -43
View File
@@ -4,11 +4,11 @@ from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising import DenoisingStage
from fastvideo.pipelines.stages.validators import StageValidators as V
from fastvideo.pipelines.stages.validators import VerificationResult
try:
from fastvideo.attention.backends.sliding_tile_attn import (
@@ -36,6 +36,7 @@ class CausalDMDDenosingStage(DenoisingStage):
def __init__(self, transformer, scheduler) -> None:
super().__init__(transformer, scheduler)
self.scheduler = FlowMatchEulerDiscreteScheduler(shift=8.0)
# KV and cross-attention cache state (initialized on first forward)
self.kv_cache1: list | None = None
self.crossattn_cache: list | None = None
@@ -70,16 +71,8 @@ class CausalDMDDenosingStage(DenoisingStage):
# Timesteps for DMD
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long).cpu()
if fastvideo_args.pipeline_config.warp_denoising_step:
logger.info("Warping timesteps...")
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
torch.tensor([0],
dtype=torch.float32)))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
logger.info("Using timesteps: %s", timesteps)
dtype=torch.long,
device=get_local_torch_device())
# Image kwargs (kept empty unless caller provides compatible args)
image_kwargs: dict = {}
@@ -269,14 +262,10 @@ class CausalDMDDenosingStage(DenoisingStage):
attn_metadata=attn_metadata,
forward_batch=batch):
# Run transformer; follow DMD stage pattern
t_expanded_noise = t_cur * torch.ones(
(latent_model_input.shape[0], 1),
device=latent_model_input.device,
dtype=torch.long)
pred_noise_btchw = self.transformer(
latent_model_input,
prompt_embeds,
t_expanded_noise,
t_expand,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=(pos_start_base + start_index) *
@@ -337,11 +326,10 @@ class CausalDMDDenosingStage(DenoisingStage):
set_forward_context(current_timestep=0,
attn_metadata=attn_metadata,
forward_batch=batch):
t_expanded_context = t_context.unsqueeze(1)
_ = self.transformer(
context_bcthw,
prompt_embeds,
t_expanded_context,
t_context,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=(pos_start_base + start_index) *
@@ -419,27 +407,3 @@ class CausalDMDDenosingStage(DenoisingStage):
False,
})
self.crossattn_cache = crossattn_cache
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify denoising stage inputs."""
result = VerificationResult()
result.add_check("latents", batch.latents,
[V.is_tensor, V.with_dims(5)])
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
result.add_check("image_embeds", batch.image_embeds, V.is_list)
result.add_check("image_latent", batch.image_latent,
V.none_or_tensor_with_dims(5))
result.add_check("num_inference_steps", batch.num_inference_steps,
V.positive_int)
result.add_check("guidance_scale", batch.guidance_scale,
V.positive_float)
result.add_check("eta", batch.eta, V.non_negative_float)
result.add_check("generator", batch.generator,
V.generator_or_list_generators)
result.add_check("do_classifier_free_guidance",
batch.do_classifier_free_guidance, V.bool_value)
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
not batch.do_classifier_free_guidance or V.list_not_empty(x))
return result
+50 -93
View File
@@ -50,63 +50,6 @@ class DecodingStage(PipelineStage):
result.add_check("output", batch.output, [V.is_tensor, V.with_dims(5)])
return result
@torch.no_grad()
def decode(self, latents: torch.Tensor,
fastvideo_args: FastVideoArgs) -> torch.Tensor:
"""
Decode latent representations into pixel space using VAE.
Args:
latents: Input latent tensor with shape (batch, channels, frames, height_latents, width_latents)
fastvideo_args: Configuration containing:
- disable_autocast: Whether to disable automatic mixed precision (default: False)
- pipeline_config.vae_precision: VAE computation precision ("fp32", "fp16", "bf16")
- pipeline_config.vae_tiling: Whether to enable VAE tiling for memory efficiency
Returns:
Decoded video tensor with shape (batch, channels, frames, height, width),
normalized to [0, 1] range and moved to CPU as float32
"""
self.vae = self.vae.to(get_local_torch_device())
latents = latents.to(get_local_torch_device())
# Setup 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
if isinstance(self.vae.scaling_factor, torch.Tensor):
latents = latents / self.vae.scaling_factor.to(
latents.device, latents.dtype)
else:
latents = latents / self.vae.scaling_factor
# Apply shifting if needed
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
latents += self.vae.shift_factor.to(latents.device,
latents.dtype)
else:
latents += self.vae.shift_factor
# Decode latents
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
if not vae_autocast_enabled:
latents = latents.to(vae_dtype)
image = self.vae.decode(latents)
# Normalize image to [0, 1] range
image = (image / 2 + 0.5).clamp(0, 1)
return image
@torch.no_grad()
def forward(
self,
@@ -116,28 +59,13 @@ class DecodingStage(PipelineStage):
"""
Decode latent representations into pixel space.
This method processes the batch through the VAE decoder, converting latent
representations to pixel-space video/images. It also optionally decodes
trajectory latents for visualization purposes.
Args:
batch: The current batch containing:
- latents: Tensor to decode (batch, channels, frames, height_latents, width_latents)
- return_trajectory_decoded (optional): Flag to decode trajectory latents
- trajectory_latents (optional): Latents at different timesteps
- trajectory_timesteps (optional): Corresponding timesteps
fastvideo_args: Configuration containing:
- output_type: "latent" to skip decoding, otherwise decode to pixels
- vae_cpu_offload: Whether to offload VAE to CPU after decoding
- model_loaded: Track VAE loading state
- model_paths: Path to VAE model if loading needed
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
Modified batch with:
- output: Decoded frames (batch, channels, frames, height, width) as CPU float32
- trajectory_decoded (if requested): List of decoded frames per timestep
The batch with decoded outputs.
"""
# load vae if not already loaded (used for memory constrained devices)
pipeline = self.pipeline() if self.pipeline else None
if not fastvideo_args.model_loaded["vae"]:
loader = VAELoader()
@@ -147,29 +75,58 @@ class DecodingStage(PipelineStage):
pipeline.add_module("vae", self.vae)
fastvideo_args.model_loaded["vae"] = True
if fastvideo_args.output_type == "latent":
frames = batch.latents
else:
frames = self.decode(batch.latents, fastvideo_args)
self.vae = self.vae.to(get_local_torch_device())
# decode trajectory latents if needed
if batch.return_trajectory_decoded:
batch.trajectory_decoded = []
assert batch.trajectory_latents is not None, "batch should have trajectory latents"
for idx in range(batch.trajectory_latents.shape[1]):
# batch.trajectory_latents is [batch_size, timesteps, channels, frames, height, width]
cur_latent = batch.trajectory_latents[:, idx, :, :, :, :]
cur_timestep = batch.trajectory_timesteps[idx]
logger.info("decoding trajectory latent for timestep: %s",
cur_timestep)
decoded_frames = self.decode(cur_latent, fastvideo_args)
batch.trajectory_decoded.append(decoded_frames.cpu().float())
latents = batch.latents
# TODO(will): remove this once we add input/output validation for stages
if latents is None:
raise ValueError("Latents must be provided")
# Skip decoding if output type is latent
if fastvideo_args.output_type == "latent":
image = latents
else:
# Setup 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
if isinstance(self.vae.scaling_factor, torch.Tensor):
latents = latents / self.vae.scaling_factor.to(
latents.device, latents.dtype)
else:
latents = latents / self.vae.scaling_factor
# Apply shifting if needed
if (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
latents += self.vae.shift_factor.to(latents.device,
latents.dtype)
else:
latents += self.vae.shift_factor
# Decode latents
with torch.autocast(device_type="cuda",
dtype=vae_dtype,
enabled=vae_autocast_enabled):
if fastvideo_args.pipeline_config.vae_tiling:
self.vae.enable_tiling()
# if fastvideo_args.vae_sp:
# self.vae.enable_parallel()
if not vae_autocast_enabled:
latents = latents.to(vae_dtype)
image = self.vae.decode(latents)
# Normalize image to [0, 1] range
image = (image / 2 + 0.5).clamp(0, 1)
# Convert to CPU float32 for compatibility
frames = frames.cpu().float()
image = image.cpu().float()
# Update batch with decoded image
batch.output = frames
batch.output = image
# Offload models if needed
if hasattr(self, 'maybe_free_model_hooks'):
+6 -74
View File
@@ -40,13 +40,6 @@ try:
except ImportError:
st_attn_available = False
try:
from fastvideo.attention.backends.vmoba import VMOBAAttentionBackend
from fastvideo.utils import is_vmoba_available
vmoba_attn_available = is_vmoba_available()
except ImportError:
vmoba_attn_available = False
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
@@ -84,7 +77,6 @@ class DenoisingStage(PipelineStage):
supported_attention_backends=(
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.VMOBA_ATTN,
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA
) # hack
)
@@ -157,8 +149,7 @@ class DenoisingStage(PipelineStage):
# Prepare image latents and embeddings for I2V generation
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert not torch.isnan(
image_embeds[0]).any(), "image_embeds contains nan"
assert torch.isnan(image_embeds[0]).sum() == 0
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
@@ -195,13 +186,11 @@ class DenoisingStage(PipelineStage):
# Get latents and embeddings
latents = batch.latents
prompt_embeds = batch.prompt_embeds
assert not torch.isnan(
prompt_embeds[0]).any(), "prompt_embeds contains nan"
assert torch.isnan(prompt_embeds[0]).sum() == 0
if batch.do_classifier_free_guidance:
neg_prompt_embeds = batch.negative_prompt_embeds
assert neg_prompt_embeds is not None
assert not torch.isnan(
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
assert torch.isnan(neg_prompt_embeds[0]).sum() == 0
# (Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
if fastvideo_args.boundary_ratio is not None:
@@ -247,10 +236,6 @@ class DenoisingStage(PipelineStage):
patch_size[2])
seq_len = int(math.ceil(seq_len / sp_world_size)) * sp_world_size
# Initialize lists for ODE trajectory
trajectory_timesteps: list[torch.Tensor] = []
trajectory_latents: list[torch.Tensor] = []
# Run denoising loop
with self.progress_bar(total=num_inference_steps) as progress_bar:
for i, t in enumerate(timesteps):
@@ -283,9 +268,6 @@ class DenoisingStage(PipelineStage):
latent_model_input = torch.cat(
[latent_model_input, batch.image_latent],
dim=1).to(target_dtype)
assert not torch.isnan(
latent_model_input).any(), "latent_model_input contains nan"
if fastvideo_args.pipeline_config.ti2v_task and batch.pil_image is not None:
timestep = torch.stack([t]).to(get_local_torch_device())
temp_ts = (mask2[0][0][:, ::2, ::2] * timestep).flatten()
@@ -298,6 +280,7 @@ class DenoisingStage(PipelineStage):
else:
t_expand = t.repeat(latent_model_input.shape[0])
assert torch.isnan(latent_model_input).sum() == 0
latent_model_input = self.scheduler.scale_model_input(
latent_model_input, t)
@@ -342,31 +325,6 @@ class DenoisingStage(PipelineStage):
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
elif (vmoba_attn_available
and self.attn_backend == VMOBAAttentionBackend):
self.attn_metadata_builder_cls = self.attn_backend.get_builder_cls(
)
if self.attn_metadata_builder_cls is not None:
self.attn_metadata_builder = self.attn_metadata_builder_cls(
)
# Prepare V-MoBA parameters from config
moba_params = fastvideo_args.moba_config.copy()
moba_params.update({
"current_timestep":
i,
"raw_latent_shape":
batch.raw_latent_shape[2:5],
"patch_size":
fastvideo_args.pipeline_config.dit_config.
patch_size,
"device":
get_local_torch_device(),
})
attn_metadata = self.attn_metadata_builder.build(
**moba_params)
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
else:
attn_metadata = None
# TODO(will): finalize the interface. vLLM uses this to
@@ -431,11 +389,6 @@ class DenoisingStage(PipelineStage):
latents = (1. - mask2[0]) * z + mask2[0] * latents
# latents = latents.unsqueeze(0)
# save trajectory latents if needed
if batch.return_trajectory_latents:
trajectory_timesteps.append(t)
trajectory_latents.append(latents)
# Update progress bar
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and
@@ -443,28 +396,9 @@ class DenoisingStage(PipelineStage):
and progress_bar is not None):
progress_bar.update()
# Gather results if using sequence parallelism
trajectory_tensor: torch.Tensor | None = None
if trajectory_latents:
trajectory_tensor = torch.stack(trajectory_latents, dim=1)
trajectory_timesteps_tensor = torch.stack(trajectory_timesteps,
dim=0)
else:
trajectory_tensor = None
trajectory_timesteps_tensor = None
# Gather results if using sequence parallelism
if sp_group:
latents = sequence_model_parallel_all_gather(latents, dim=2)
if batch.return_trajectory_latents:
trajectory_tensor = trajectory_tensor.to(
get_local_torch_device())
trajectory_tensor = sequence_model_parallel_all_gather(
trajectory_tensor, dim=3)
if trajectory_tensor is not None and trajectory_timesteps_tensor is not None:
batch.trajectory_timesteps = trajectory_timesteps_tensor.cpu()
batch.trajectory_latents = trajectory_tensor.cpu()
# Update batch with final latents
batch.latents = latents
@@ -803,8 +737,7 @@ class DmdDenoisingStage(DenoisingStage):
video_raw_latent_shape = latents.shape
prompt_embeds = batch.prompt_embeds
assert not torch.isnan(
prompt_embeds[0]).any(), "prompt_embeds contains nan"
assert torch.isnan(prompt_embeds[0]).sum() == 0
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long,
@@ -843,8 +776,7 @@ class DmdDenoisingStage(DenoisingStage):
batch.image_latent.permute(0, 2, 1, 3, 4)
],
dim=2).to(target_dtype)
assert not torch.isnan(
latent_model_input).any(), "latent_model_input contains nan"
assert torch.isnan(latent_model_input).sum() == 0
# Prepare inputs for transformer
t_expand = t.repeat(latent_model_input.shape[0])
-14
View File
@@ -159,20 +159,6 @@ class CudaPlatformBase(Platform):
str(e))
raise ImportError(
"Video Sparse Attention backend is not installed. ") from e
elif selected_backend == AttentionBackendEnum.VMOBA_ATTN:
try:
from csrc.attn.vmoba_attn.vmoba import ( # noqa: F401
moba_attn_varlen)
from fastvideo.attention.backends.vmoba import ( # noqa: F401
VMOBAAttentionBackend)
logger.info("Using Video MOBA Attention backend.")
return "fastvideo.attention.backends.vmoba.VMOBAAttentionBackend"
except ImportError as e:
logger.error(
"Failed to import Video MoBA Attention backend: %s", str(e))
raise ImportError(
"Video MoBA Attention backend is not installed. ") from e
elif selected_backend == AttentionBackendEnum.TORCH_SDPA:
logger.info("Using Torch SDPA backend.")
return "fastvideo.attention.backends.sdpa.SDPABackend"
-1
View File
@@ -19,7 +19,6 @@ class AttentionBackendEnum(enum.Enum):
TORCH_SDPA = enum.auto()
SAGE_ATTN = enum.auto()
VIDEO_SPARSE_ATTN = enum.auto()
VMOBA_ATTN = enum.auto()
NO_ATTENTION = enum.auto()
-108
View File
@@ -1,108 +0,0 @@
import os
from pathlib import Path
import pyarrow as pa
import pyarrow.parquet as pq
from fastvideo.dataset.dataloader.parquet_io import (
ParquetDatasetWriter,
records_to_table,
)
def test_records_to_table_types():
schema = pa.schema([
pa.field("id", pa.string()),
pa.field("vae_latent_bytes", pa.binary()),
pa.field("vae_latent_shape", pa.list_(pa.int64())),
pa.field("duration_sec", pa.float64()),
pa.field("width", pa.int64()),
])
records = [{
"id": "a",
"vae_latent_bytes": b"\x00\x01",
"vae_latent_shape": [1, 2, 3],
"duration_sec": 1.5,
"width": 640,
}]
table = records_to_table(records, schema)
assert table.schema == schema
assert table.num_rows == 1
cols = {name: table.column(name).to_pylist()[0] for name in schema.names}
assert cols["id"] == "a"
assert isinstance(cols["vae_latent_bytes"], (bytes, bytearray))
assert cols["vae_latent_shape"] == [1, 2, 3]
assert abs(cols["duration_sec"] - 1.5) < 1e-6
assert cols["width"] == 640
def test_writer_flush_and_remainder(tmp_path: Path):
schema = pa.schema([pa.field("id", pa.string())])
records = [{"id": str(i)} for i in range(25)]
table = records_to_table(records, schema)
out_dir = tmp_path / "out"
writer = ParquetDatasetWriter(str(out_dir), samples_per_file=10)
writer.append_table(table)
written = writer.flush(num_workers=1)
assert written == 20
files = sorted(out_dir.rglob("*.parquet"))
assert len(files) == 2
total_rows = sum(pq.read_table(str(f)).num_rows for f in files)
assert total_rows == 20
# Append remainder to complete another chunk
extra = records_to_table([{"id": str(i)} for i in range(5)], schema)
writer.append_table(extra)
written2 = writer.flush(num_workers=1)
assert written2 == 10
files2 = sorted(out_dir.rglob("*.parquet"))
assert len(files2) == 3
total_rows2 = sum(pq.read_table(str(f)).num_rows for f in files2)
assert total_rows2 == 30
def test_writer_flush_write_remainder(tmp_path: Path):
schema = pa.schema([pa.field("id", pa.string())])
# 25 rows, 10 per file => 2 full files + 1 remainder(5)
records = [{"id": str(i)} for i in range(25)]
table = records_to_table(records, schema)
out_dir = tmp_path / "out_last"
writer = ParquetDatasetWriter(str(out_dir), samples_per_file=10)
writer.append_table(table)
# First flush writes 20
written1 = writer.flush(num_workers=1)
assert written1 == 20
# Final flush with remainder
written2 = writer.flush(num_workers=1, write_remainder=True)
assert written2 == 5
files = sorted(out_dir.rglob("*.parquet"))
assert len(files) == 3
total_rows = sum(pq.read_table(str(f)).num_rows for f in files)
assert total_rows == 25
def test_writer_parallel_workers(tmp_path: Path):
schema = pa.schema([pa.field("id", pa.string())])
# 40 rows, 10 per file => 4 files
records = [{"id": str(i)} for i in range(40)]
table = records_to_table(records, schema)
out_dir = tmp_path / "out_parallel"
writer = ParquetDatasetWriter(str(out_dir), samples_per_file=10)
writer.append_table(table)
written = writer.flush(num_workers=2)
assert written == 40
# Ensure files exist under worker subdirs
worker_dirs = [p for p in out_dir.iterdir() if p.is_dir() and p.name.startswith("worker_")]
assert len(worker_dirs) >= 1
files = sorted(out_dir.rglob("*.parquet"))
assert len(files) == 4
total_rows = sum(pq.read_table(str(f)).num_rows for f in files)
assert total_rows == 40
@@ -1,58 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import os
import subprocess
from pathlib import Path
def test_inference_vmoba():
"""Test FastVideo VMOBA_ATTN inference pipeline"""
num_gpus = "1"
model_base = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
output_dir = Path("outputs_video/vmoba_1.3B/")
moba_config = "fastvideo/configs/backend/vmoba/wan_1.3B_77_480_832.json"
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = "VMOBA_ATTN"
cmd = [
"fastvideo", "generate",
"--model-path", model_base,
"--sp-size", num_gpus,
"--tp-size", "1",
"--num-gpus", num_gpus,
"--dit-cpu-offload", "False",
"--vae-cpu-offload", "False",
"--text-encoder-cpu-offload", "True",
"--pin-cpu-memory", "False",
"--height", "480",
"--width", "832",
"--num-frames", "77",
"--num-inference-steps", "50",
"--moba-config-path", moba_config,
"--fps", "16",
"--guidance-scale", "6.0",
"--flow-shift", "8.0",
"--prompt", "A majestic lion strides across the golden savanna, its powerful frame glistening under the warm afternoon sun. The tall grass ripples gently in the breeze, enhancing the lion's commanding presence. The tone is vibrant, embodying the raw energy of the wild. Low angle, steady tracking shot, cinematic.",
"--negative-prompt", (
"Bright tones, overexposed, static, blurred details, subtitles, style, "
"works, paintings, images, static, overall gray, worst quality, low quality, "
"JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, "
"poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, "
"still picture, messy background, three legs, many people in the background, walking backwards"
),
"--seed", "1024",
"--output-path", str(output_dir),
]
subprocess.run(cmd, check=True)
assert output_dir.exists(), f"Output directory {output_dir} does not exist"
video_files = list(output_dir.glob("*.mp4"))
assert len(video_files) > 0, "No video files were generated"
for video_file in video_files:
assert video_file.stat().st_size > 0, f"Video file {video_file} is empty"
if __name__ == "__main__":
test_inference_vmoba()
+1 -9
View File
@@ -102,18 +102,10 @@ def run_precision_tests_STA():
def run_precision_tests_VSA():
run_test("python csrc/attn/tests/test_vsa.py")
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_precision_tests_vmoba():
run_test("pytest csrc/attn/vmoba_attn/tests/test_vmoba_attn.py")
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_inference_tests_vmoba():
run_test('python fastvideo/tests/inference/vmoba/test_vmoba_inference.py')
@app.function(gpu="L40S:1", image=image, timeout=3600)
def run_inference_lora_tests():
run_test("pytest ./fastvideo/tests/inference/lora/test_lora_inference_similarity.py -vs")
@app.function(gpu="L40S:2", image=image, timeout=900)
def run_distill_dmd_tests():
run_test("pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
run_test("pytest ./fastvideo/tests/training/distill/test_distill_dmd.py -vs")
@@ -1,67 +0,0 @@
from pathlib import Path
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
from fastvideo.workflow.preprocess.components import ParquetDatasetSaver
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
def _simple_record_creator(batch: PreprocessBatch) -> list[dict]:
# batch.latents will be converted to numpy by the saver before this call
assert isinstance(batch.latents, np.ndarray)
num = len(batch.video_file_name)
records = []
for i in range(num):
arr = batch.latents[i]
records.append({
"id": batch.video_file_name[i],
"data_bytes": arr.tobytes(),
"data_shape": list(arr.shape),
})
return records
def test_parquet_dataset_saver_flush_and_last(tmp_path: Path):
# Schema for the simple record creator
schema = pa.schema([
pa.field("id", pa.string()),
pa.field("data_bytes", pa.binary()),
pa.field("data_shape", pa.list_(pa.int64())),
])
B = 5
# Build a minimal PreprocessBatch
batch = PreprocessBatch(
data_type="video",
latents=torch.randn(B, 2),
prompt_embeds=[torch.randn(B, 1, 1)],
prompt_attention_mask=[torch.ones(B, 1)],
)
batch.video_file_name = [f"vid_{i}" for i in range(B)]
saver = ParquetDatasetSaver(
flush_frequency=10, # higher than B to avoid auto-flush
samples_per_file=3,
schema=schema,
record_creator=_simple_record_creator,
)
out_dir = tmp_path / "saver_out"
saver.save_and_write_parquet_batch(batch, str(out_dir))
# First flush: should write one full file (3 rows), keep 2 in buffer
saver.flush_tables(str(out_dir))
files = sorted(out_dir.rglob("*.parquet"))
assert len(files) == 1
assert pq.read_table(str(files[0])).num_rows == 3
# Final flush: write remainder 2 rows
saver.flush_last(str(out_dir))
files2 = sorted(out_dir.rglob("*.parquet"))
assert len(files2) == 2
total = sum(pq.read_table(str(f)).num_rows for f in files2)
assert total == 5
+9 -18
View File
@@ -36,9 +36,8 @@ from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing)
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases, count_trainable,
get_scheduler, load_distillation_checkpoint, save_distillation_checkpoint,
shift_timestep)
clip_grad_norm_while_handling_failing_dtensor_cases, get_scheduler,
load_distillation_checkpoint, save_distillation_checkpoint, shift_timestep)
from fastvideo.utils import is_vsa_available, set_random_seed
import wandb # isort: skip
@@ -69,18 +68,11 @@ class DistillationPipeline(TrainingPipeline):
current_trainstep: int
video_latent_shape: tuple[int, ...]
video_latent_shape_sp: tuple[int, ...]
real_score_transformer: torch.nn.Module
fake_score_transformer: torch.nn.Module
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
raise RuntimeError(
"create_pipeline_stages should not be called for training pipeline")
def set_trainable(self) -> None:
super().set_trainable()
self.modules["real_score_transformer"].requires_grad_(False)
self.modules["vae"].requires_grad_(False)
def initialize_training_pipeline(self, training_args: TrainingArgs):
"""Initialize the distillation training pipeline with multiple models."""
logger.info("Initializing distillation pipeline...")
@@ -89,6 +81,8 @@ class DistillationPipeline(TrainingPipeline):
self.noise_scheduler = self.get_module("scheduler")
self.vae = self.get_module("vae")
self.vae.requires_grad_(False)
self.timestep_shift = self.training_args.pipeline_config.flow_shift
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
shift=self.timestep_shift)
@@ -96,7 +90,10 @@ class DistillationPipeline(TrainingPipeline):
# self.transformer is the generator model
self.real_score_transformer = self.get_module("real_score_transformer")
self.fake_score_transformer = self.get_module("fake_score_transformer")
self.real_score_transformer.requires_grad_(False)
self.real_score_transformer.eval()
self.fake_score_transformer.requires_grad_(True)
self.fake_score_transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
@@ -170,7 +167,9 @@ class DistillationPipeline(TrainingPipeline):
def _prepare_distillation(self,
training_batch: TrainingBatch) -> TrainingBatch:
"""Prepare training environment for distillation."""
self.transformer.requires_grad_(True)
self.transformer.train()
self.fake_score_transformer.requires_grad_(True)
self.fake_score_transformer.train()
return training_batch
@@ -896,14 +895,6 @@ class DistillationPipeline(TrainingPipeline):
else:
set_random_seed(seed + self.global_rank)
# Check trainable params
num_trainable_generator = round(
count_trainable(self.transformer) / 1e9, 3)
num_trainable_critic = round(
count_trainable(self.fake_score_transformer) / 1e9, 3)
logger.info(
"rank: %s: # of trainable params in generator: %sB, # of trainable params in critic: %sB",
self.global_rank, num_trainable_generator, num_trainable_critic)
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
+12 -28
View File
@@ -22,7 +22,6 @@ from tqdm.auto import tqdm
import fastvideo.envs as envs
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadataBuilder)
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import build_parquet_map_style_dataloader
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
@@ -39,20 +38,22 @@ from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing)
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases,
compute_density_for_timestep_sampling, count_trainable, get_scheduler,
get_sigmas, load_checkpoint, normalize_dit_input, save_checkpoint,
compute_density_for_timestep_sampling, get_scheduler, get_sigmas,
load_checkpoint, normalize_dit_input, save_checkpoint,
shard_latents_across_sp)
from fastvideo.utils import (is_vmoba_available, is_vsa_available,
set_random_seed, shallow_asdict)
from fastvideo.utils import is_vsa_available, set_random_seed, shallow_asdict
import wandb # isort: skip
vsa_available = is_vsa_available()
vmoba_available = is_vmoba_available()
logger = init_logger(__name__)
def _get_trainable_params(model: torch.nn.Module) -> int:
return sum(p.numel() for p in model.parameters() if p.requires_grad)
class TrainingPipeline(LoRAPipeline, ABC):
"""
A pipeline for training a model. All training pipelines should inherit from this class.
@@ -112,7 +113,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
enable_gradient_checkpointing_type)
noise_scheduler = self.modules["scheduler"]
# Set grads for proper modules based on the training mode (Distill, LoRA, etc.)
self.set_trainable()
params_to_optimize = self.transformer.parameters()
params_to_optimize = list(
@@ -272,20 +272,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
patch_size=patch_size,
VSA_sparsity=current_vsa_sparsity,
device=get_local_torch_device())
elif vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
moba_params = self.training_args.moba_config.copy()
moba_params.update({
"current_timestep":
training_batch.timesteps,
"raw_latent_shape":
training_batch.raw_latent_shape[2:5],
"patch_size":
self.training_args.pipeline_config.dit_config.patch_size,
"device":
get_local_torch_device(),
})
training_batch.attn_metadata = VideoMobaAttentionMetadataBuilder(
).build(**moba_params)
else:
training_batch.attn_metadata = None
@@ -310,7 +296,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
def _transformer_forward_and_compute_loss(
self, training_batch: TrainingBatch) -> TrainingBatch:
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN" or vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
assert training_batch.attn_metadata is not None
else:
assert training_batch.attn_metadata is None
@@ -431,7 +417,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
local_main_process_only=False)
if not self.post_init_called:
self.post_init()
num_trainable_params = count_trainable(self.transformer)
num_trainable_params = _get_trainable_params(self.transformer)
logger.info("Starting training with %s B trainable parameters",
round(num_trainable_params / 1e9, 3))
@@ -476,9 +462,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
current_decay_times = min(step // vsa_decay_interval_steps,
vsa_sparsity // vsa_decay_rate)
current_vsa_sparsity = current_decay_times * vsa_decay_rate
elif vmoba_available:
# TODO: add vmoba sparsity scheduling here
pass
else:
current_vsa_sparsity = 0.0
@@ -523,7 +506,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
self._log_validation(self.transformer, self.training_args, step)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
trainable_params = round(
count_trainable(self.transformer) / 1e9, 3)
_get_trainable_params(self.transformer) / 1e9, 3)
logger.info(
"GPU memory usage after validation: %s MB, trainable params: %sB",
gpu_memory_usage, trainable_params)
@@ -559,7 +542,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
logger.info(" Total optimization steps = %s",
self.training_args.max_train_steps)
logger.info(" Total training parameters per FSDP shard = %s B",
round(count_trainable(self.transformer) / 1e9, 3))
round(_get_trainable_params(self.transformer) / 1e9, 3))
# print dtype
logger.info(" Master weight dtype: %s",
self.transformer.parameters().__next__().dtype)
@@ -627,6 +610,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
validation_dataloader = DataLoader(validation_dataset,
batch_size=None,
num_workers=0)
transformer.eval()
validation_steps = training_args.validation_sampling_steps.split(",")
-4
View File
@@ -1278,7 +1278,3 @@ def get_scheduler(
num_warmup_steps=num_warmup_steps,
num_training_steps=num_training_steps,
last_epoch=last_epoch)
def count_trainable(model: torch.nn.Module) -> int:
return sum(p.numel() for p in model.parameters() if p.requires_grad)
-29
View File
@@ -23,14 +23,10 @@ from typing import Any, TypeVar, cast
import cloudpickle
import filelock
import imageio
import numpy as np
import torch
import torchvision
import yaml
from diffusers.loaders.lora_base import (
_best_guess_weight_name) # watch out for potetential removal from diffusers
from einops import rearrange
from huggingface_hub import snapshot_download
from remote_pdb import RemotePdb
from torch.distributed.fsdp import MixedPrecisionPolicy
@@ -818,17 +814,6 @@ def is_vsa_available() -> bool:
return importlib.util.find_spec("vsa") is not None
@lru_cache(maxsize=1)
def is_vmoba_available() -> bool:
if importlib.util.find_spec("csrc.attn.vmoba_attn.vmoba") is None:
return False
try:
import flash_attn
return flash_attn.__version__ >= "2.7.4"
except Exception:
return False
# adapted from: https://github.com/Wan-Video/Wan2.2/blob/main/wan/utils/utils.py
def masks_like(tensor,
zero=False,
@@ -890,17 +875,3 @@ def best_output_size(w, h, dw, dh, expected_area):
return ow1, oh1
else:
return ow2, oh2
def save_decoded_latents_as_video(decoded_latents: list[torch.Tensor],
output_path: str, fps: int):
# Process outputs
videos = rearrange(decoded_latents, "b c t h w -> t b c h w")
frames = []
for x in videos:
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))
os.makedirs(os.path.dirname(output_path), exist_ok=True)
imageio.mimsave(output_path, frames, fps=fps, format="mp4")
+157 -33
View File
@@ -1,19 +1,20 @@
import dataclasses
import gc
import multiprocessing
import os
import random
from collections.abc import Callable
from concurrent.futures import ProcessPoolExecutor
from typing import Any
import numpy as np
import pyarrow as pa
import pyarrow.parquet as pq
import torch
from datasets import Dataset, Video, load_dataset
from fastvideo.configs.configs import (DatasetType, PreprocessConfig,
VideoLoaderType)
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
records_to_table)
from fastvideo.distributed.parallel_state import get_world_rank, get_world_size
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
@@ -154,17 +155,31 @@ class VideoForwardBatchBuilder:
class ParquetDatasetSaver:
"""Component for saving and writing Parquet datasets using shared parquet_io."""
"""Component for saving and writing Parquet datasets"""
def __init__(self, flush_frequency: int, samples_per_file: int,
schema: pa.Schema,
record_creator: Callable[..., list[dict[str, Any]]]):
def __init__(self,
flush_frequency: int,
samples_per_file: int,
schema_fields: list[str],
record_creator: Callable[..., list[dict[str, Any]]],
file_writer_fn: Callable | None = None):
"""
Initialize ParquetDatasetSaver
Args:
schema_fields: schema fields list
record_creator: Function for creating records
file_writer_fn: Function for writing records to files, uses default implementation if None
"""
self.flush_frequency = flush_frequency
self.samples_per_file = samples_per_file
self.schema = schema
self.schema_fields = schema_fields
self.create_records_from_batch = record_creator
self.file_writer_fn: Callable[
[tuple], int] = file_writer_fn or self._default_file_writer_fn
self.all_tables: list[pa.Table] = []
self.num_processed_samples: int = 0
self._writer: ParquetDatasetWriter | None = None
self.num_saved_files: int = 0
def save_and_write_parquet_batch(
self,
@@ -212,12 +227,11 @@ class ParquetDatasetSaver:
if batch_data:
self.num_processed_samples += len(batch_data)
table = records_to_table(batch_data, self.schema)
if self._writer is None:
os.makedirs(output_dir, exist_ok=True)
self._writer = ParquetDatasetWriter(
out_dir=output_dir, samples_per_file=self.samples_per_file)
self._writer.append_table(table)
# Convert batch data to PyArrow arrays
table = self._convert_batch_to_pyarrow_table(batch_data)
# Store the table in a list for later processing
self.all_tables.append(table)
logger.debug("Collected batch with %s samples", len(table))
# If flush is needed
@@ -247,35 +261,145 @@ class ParquetDatasetSaver:
def _convert_batch_to_pyarrow_table(self,
batch_data: list[dict]) -> pa.Table:
# Deprecated path, kept for backward compatibility if needed.
return records_to_table(batch_data, self.schema)
"""Convert batch data to PyArrow table"""
arrays = []
def flush_tables(self, output_dir: str, write_remainder: bool = False):
"""Flush buffered records to disk.
for field in self.schema_fields:
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
elif field in ['width', 'height', 'num_frames']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.int32()))
elif field in ['duration_sec', 'fps']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.float32()))
else:
arrays.append(pa.array([record[field]
for record in batch_data]))
Args:
output_dir: Directory where parquet files are written. Kept for API
symmetry (writer already configured with this path).
write_remainder: If True, also write any leftover rows smaller than
``samples_per_file`` as a final small file. Useful for the last flush.
"""
if self._writer is None:
return pa.Table.from_arrays(arrays, names=self.schema_fields)
def flush_tables(self, output_dir: str):
"""Flush collected tables to disk"""
if not hasattr(self, 'all_tables') or not self.all_tables:
return
_ = self._writer.flush(write_remainder=write_remainder)
# Reset processed sample count modulo samples_per_file
remainder = self.num_processed_samples % self.samples_per_file
self.num_processed_samples = 0 if write_remainder else remainder
def flush_last(self, output_dir: str):
"""Flush and write any remaining rows (final flush)."""
self.flush_tables(output_dir, write_remainder=True)
logger.debug("Combining %d batches...", len(self.all_tables))
combined_table = pa.concat_tables(self.all_tables)
assert len(combined_table) == self.num_processed_samples
logger.debug("Total samples collected: %d", len(combined_table))
# Calculate total number of chunks needed, putting remainder into self.all_tables
total_files = max(self.num_processed_samples // self.samples_per_file,
1)
logger.debug("Fixed samples per parquet file: %d",
self.samples_per_file)
logger.debug("Total number of parquet files: %d", total_files)
logger.debug(
"Total samples to be processed: %d (putting %d samples into self.all_tables)",
total_files * self.samples_per_file,
self.num_processed_samples % self.samples_per_file)
# Split work among processes
num_workers = int(min(multiprocessing.cpu_count(), total_files))
files_per_worker = (total_files + num_workers - 1) // num_workers
logger.debug("Using %d workers to process %d files", num_workers,
total_files)
logger.debug("Files per worker: %s", files_per_worker)
# Prepare work ranges
work_ranges = []
for i in range(num_workers):
start_idx = i * files_per_worker
end_idx = min((i + 1) * files_per_worker, total_files)
if start_idx < total_files:
work_ranges.append((start_idx, end_idx, combined_table, i,
output_dir, self.samples_per_file))
total_written = 0
failed_ranges = []
with ProcessPoolExecutor(max_workers=num_workers) as executor:
futures = {
executor.submit(self.file_writer_fn, work_range): work_range
for work_range in work_ranges
}
for future in futures:
try:
written = future.result()
total_written += written
logger.info("Processed file with %s samples", written)
except Exception as e:
work_range = futures[future]
failed_ranges.append(work_range)
logger.error("Failed to process range %s-%s: %s",
work_range[0], work_range[1], str(e))
# Retry failed ranges sequentially
if failed_ranges:
logger.warning("Retrying %s failed ranges sequentially",
len(failed_ranges))
for work_range in failed_ranges:
try:
total_written += self.file_writer_fn(work_range)
except Exception as e:
logger.error(
"Failed to process range %s-%s after retry: %s",
work_range[0], work_range[1], str(e))
self.num_saved_files += total_files
# Clear tables list
self.all_tables = []
if self.num_processed_samples > self.samples_per_file:
saved_samples = total_files * self.samples_per_file
self.all_tables.append(combined_table.slice(saved_samples))
self.num_processed_samples -= saved_samples
else:
self.num_processed_samples = 0
del combined_table
gc.collect()
def clean_up(self) -> None:
"""Clean up all tables"""
self._writer = None
self.all_tables = []
self.num_processed_samples = 0
self.num_saved_files = 0
gc.collect()
def _default_file_writer_fn(self, args_tuple: tuple) -> int:
"""Default chunk processing implementation"""
start_idx, end_idx, combined_table, worker_id, output_dir, samples_per_file = args_tuple
written_count = 0
for file_idx in range(start_idx, end_idx):
start_row = file_idx * samples_per_file
end_row = min(start_row + samples_per_file, len(combined_table))
if start_row >= len(combined_table):
break
chunk_table = combined_table.slice(start_row, end_row - start_row)
# Write to file
output_file = os.path.join(
output_dir,
f"chunk_{file_idx + self.num_saved_files:06d}.parquet")
pq.write_table(chunk_table, output_file)
written_count += len(chunk_table)
return written_count
def build_dataset(preprocess_config: PreprocessConfig, split: str,
validator: Callable[[dict[str, Any]], bool]) -> Dataset:
@@ -85,16 +85,16 @@ class PreprocessWorkflow(WorkflowBase):
# record creator
if self.fastvideo_args.workload_type == WorkloadType.I2V:
record_creator = i2v_record_creator
schema = pyarrow_schema_i2v
schema_fields = [f.name for f in pyarrow_schema_i2v]
else:
record_creator = basic_t2v_record_creator
schema = pyarrow_schema_t2v
schema_fields = [f.name for f in pyarrow_schema_t2v]
processed_dataset_saver = ParquetDatasetSaver(
flush_frequency=self.fastvideo_args.preprocess_config.
flush_frequency,
samples_per_file=self.fastvideo_args.preprocess_config.
samples_per_file,
schema=schema,
schema_fields=schema_fields,
record_creator=record_creator,
)
self.add_component("processed_dataset_saver", processed_dataset_saver)
@@ -1,28 +0,0 @@
#!/bin/bash
num_gpus=1
export FASTVIDEO_ATTENTION_BACKEND=VMOBA_ATTN
export MODEL_BASE=FastVideo/Wan2.1-T2V-1.3B-Diffusers
# export MODEL_BASE=hunyuanvideo-community/HunyuanVideo
# You can either use --prompt or --prompt-txt, but not both.
fastvideo generate \
--model-path $MODEL_BASE \
--sp-size $num_gpus \
--tp-size 1 \
--num-gpus $num_gpus \
--dit-cpu-offload False \
--vae-cpu-offload False \
--text-encoder-cpu-offload True \
--pin-cpu-memory False \
--height 480 \
--width 832 \
--num-frames 77 \
--num-inference-steps 50 \
--moba-config-path fastvideo/configs/backend/vmoba/wan_1.3B_77_480_832.json \
--fps 16 \
--guidance-scale 6.0 \
--flow-shift 8.0 \
--prompt-txt assets/prompt.txt \
--negative-prompt "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards" \
--seed 1024 \
--output-path outputs_video/
+21
View File
@@ -0,0 +1,21 @@
#!/bin/bash
# Create output directory if it doesn't exist
mkdir -p preprocess_output
# Launch 8 jobs, one for each node
# Each node processes 8 consecutive files (64 total files / 8 nodes = 8 files per node)
for node_id in {0..1}; do
# Calculate the starting file number for this node
start_file=$((node_id * 8 + 1))
echo "Launching node $node_id with files v2m_${start_file}.txt to v2m_$((start_file + 7)).txt"
echo "sbatch --job-name=ode-${node_id} --output=preprocess_output/preprocess-node-${node_id}.out --error=preprocess_output/preprocess-node-${node_id}.err /mnt/weka/home/hao.zhang/wl/kev/preproc/slurms/syn.slurm $start_file $node_id"
sbatch --job-name=ode-${node_id} \
--output=preprocess_output/preprocess-node-${node_id}.out \
--error=preprocess_output/preprocess-node-${node_id}.err \
/mnt/weka/home/hao.zhang/wl/kev/preproc/slurms/syn.slurm $start_file $node_id
done
echo "All 8 nodes launched successfully!"
+21
View File
@@ -0,0 +1,21 @@
#!/bin/bash
# Create output directory if it doesn't exist
mkdir -p preprocess_output_text
# Launch 8 jobs, one for each node
# Each node processes 8 consecutive files (64 total files / 8 nodes = 8 files per node)
for node_id in {0..1}; do
# Calculate the starting file number for this node
start_file=$((node_id * 8 + 1))
echo "Launching text-only node $node_id with files v2m_${start_file}.txt to v2m_$((start_file + 7)).txt"
echo "sbatch --job-name=text-${node_id} --output=preprocess_output_text/preprocess-text-node-${node_id}.out --error=preprocess_output_text/preprocess-text-node-${node_id}.err /mnt/weka/home/hao.zhang/matthew/FastVideo/scripts/preprocess/syn_text.slurm $start_file $node_id"
sbatch --job-name=text-${node_id} \
--output=preprocess_output_text/preprocess-text-node-${node_id}.out \
--error=preprocess_output_text/preprocess-text-node-${node_id}.err \
/mnt/weka/home/hao.zhang/matthew/FastVideo/scripts/preprocess/syn_text.slurm $start_file $node_id
done
echo "All 8 text-only nodes launched successfully!"
+88
View File
@@ -0,0 +1,88 @@
#!/bin/bash
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=1
#SBATCH --ntasks-per-node=8
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=16
#SBATCH --mem=960G
#SBATCH --exclusive
#SBATCH --time=72:00:00
# conda init
# source ~/conda/miniconda/bin/activate
# PYTHON_VIRTUAL_ENVIRONMENT=fastvideo-train-yq
# conda activate $PYTHON_VIRTUAL_ENVIRONMENT
nvidia-smi
nvidia-smi --query-compute-apps=pid,process_name,used_memory --format=csv
echo " "
echo " Number of nodes:= " $SLURM_JOB_NUM_NODES
echo " GPUs per node:= " $SLURM_JOB_GPUS
echo " Running on multiple nodes/GPU devices"
echo ""
echo " Run started at:- "
date
# Accept parameters from launch script
START_FILE=${1:-1} # Starting file number for this node
NODE_ID=${2:-0} # Node identifier (0-7)
num_gpus=1
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
# Start port number - we'll increment for each job
base_port=$((29603 + NODE_ID * 100)) # Different port range per node
# Create an array of CUDA device IDs
gpu_ids=(0 1 2 3 4 5 6 7)
GPU_NUM=1
MODEL_TYPE="wan"
echo "NODE_ID: $NODE_ID"
echo "START_FILE: $START_FILE"
echo "Base port for this node: $base_port"
# Run 8 parallel preprocessing jobs on this node
for i in {1..8}; do
# Calculate port for this job
port=$((base_port + i))
# Get GPU ID using modulo to cycle through available GPUs
gpu=${gpu_ids[((i-1))]}
# Calculate which file this GPU should process
file_num=$((START_FILE + i - 1))
DATA_MERGE_PATH="/mnt/weka/home/hao.zhang/wl/kev/preproc/prompts/v2m_${file_num}.txt"
# Create unique output directory based on node and GPU
OUTPUT_DIR="data/test-ode-preprocessing/Node_${NODE_ID}_GPU_${i}_File_${file_num}"
start_cpu=$(( (i-1)*2 )) # Reduced CPU allocation for 8 nodes
end_cpu=$(( start_cpu+1 ))
echo "Starting GPU $gpu processing file v2m_${file_num}.txt on port $port, output: $OUTPUT_DIR"
# Run the preprocessing command in background
CUDA_VISIBLE_DEVICES=$gpu taskset -c ${start_cpu}-${end_cpu} torchrun --nnodes=1 --nproc_per_node=$GPU_NUM --master_port $port \
/mnt/weka/home/hao.zhang/wl/kev/FastVideo/fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_BASE \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 2 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "ode_trajectory" &
done
# Wait for all jobs on this node to complete
wait
echo "All processing blocks completed!"
+89
View File
@@ -0,0 +1,89 @@
#!/bin/bash
#SBATCH --partition=main
#SBATCH --qos=hao
#SBATCH --nodes=1
#SBATCH --ntasks-per-node=8
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=16
#SBATCH --mem=960G
#SBATCH --exclusive
#SBATCH --time=72:00:00
# conda init
# source ~/conda/miniconda/bin/activate
# PYTHON_VIRTUAL_ENVIRONMENT=fastvideo-train-yq
# conda activate $PYTHON_VIRTUAL_ENVIRONMENT
nvidia-smi
nvidia-smi --query-compute-apps=pid,process_name,used_memory --format=csv
echo " "
echo " Number of nodes:= " $SLURM_JOB_NUM_NODES
echo " GPUs per node:= " $SLURM_JOB_GPUS
echo " Running on multiple nodes/GPU devices for TEXT-ONLY preprocessing"
echo ""
echo " Run started at:- "
date
# Accept parameters from launch script
START_FILE=${1:-1} # Starting file number for this node
NODE_ID=${2:-0} # Node identifier (0-7)
num_gpus=1
export MODEL_BASE=Wan-AI/Wan2.1-T2V-1.3B-Diffusers
# Start port number - we'll increment for each job
base_port=$((29603 + NODE_ID * 100)) # Different port range per node
# Create an array of CUDA device IDs
gpu_ids=(0 1 2 3 4 5 6 7)
GPU_NUM=1
MODEL_TYPE="wan"
echo "NODE_ID: $NODE_ID"
echo "START_FILE: $START_FILE"
echo "Base port for this node: $base_port"
echo "Processing TEXT-ONLY data"
# Run 8 parallel preprocessing jobs on this node
for i in {1..8}; do
# Calculate port for this job
port=$((base_port + i))
# Get GPU ID using modulo to cycle through available GPUs
gpu=${gpu_ids[((i-1))]}
# Calculate which file this GPU should process
file_num=$((START_FILE + i - 1))
DATA_MERGE_PATH="/mnt/weka/home/hao.zhang/wl/kev/preproc/prompts/v2m_${file_num}.txt"
# Create unique output directory based on node and GPU for text-only processing
OUTPUT_DIR="data/test-text-preprocessing/Node_${NODE_ID}_GPU_${i}_File_${file_num}"
start_cpu=$(( (i-1)*2 )) # Reduced CPU allocation for 8 nodes
end_cpu=$(( start_cpu+1 ))
echo "Starting GPU $gpu processing text-only file v2m_${file_num}.txt on port $port, output: $OUTPUT_DIR"
# Run the text-only preprocessing command in background
CUDA_VISIBLE_DEVICES=$gpu taskset -c ${start_cpu}-${end_cpu} torchrun --nnodes=1 --nproc_per_node=$GPU_NUM --master_port $port \
/mnt/weka/home/hao.zhang/matthew/FastVideo/fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_BASE \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 2 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "text" &
done
# Wait for all jobs on this node to complete
wait
echo "All text-only processing blocks completed!"