Compare commits

..
Author SHA1 Message Date
SolitaryThinker 352e3c31fe tests 2025-09-10 01:57:42 +00:00
SolitaryThinker 4f5e79c41f update 2025-09-10 01:52:21 +00:00
SolitaryThinker a2d303b067 fix rebase 2025-09-09 23:30:19 +00:00
SolitaryThinker e07111b0de update parquet handling 2025-09-09 23:29:37 +00:00
SolitaryThinker aa49d2a5c8 fix num_inferenc_steps 2025-09-09 23:29:37 +00:00
SolitaryThinker b337d03e82 disable trajectory deocding 2025-09-09 23:29:37 +00:00
SolitaryThinker b92da9e912 hack to get it running 2025-09-09 23:29:36 +00:00
SolitaryThinkerandkevin314 d615271814 add kevin as coauthor
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
2025-09-09 23:29:36 +00:00
SolitaryThinker 720cfe39ca update 2025-09-09 23:29:36 +00:00
SolitaryThinker 03da4b1cdc update 2025-09-09 23:29:36 +00:00
SolitaryThinker ce52c3e87e rename 2025-09-09 23:29:35 +00:00
SolitaryThinker 7b7a895e77 checkpoint 2025-09-09 23:29:35 +00:00
SolitaryThinker 4744ec2c0b checkpoint 2025-09-09 23:29:33 +00:00
William LinandRandNMR73 ac11127397 [Self-forcing] [1/n] Handle extra dim in time embedding and add timestep warping (#792)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
2025-09-09 02:52:02 -07:00
Eric LiangandEricLiang e028dcc7c0 [Backend][Vmoba] Add implementation of VMoba (#778)
Co-authored-by: EricLiang <https://github.com/EricLina>
2025-09-08 23:53:25 -07:00
Wenxuan Tanandgemini-code-assist[bot] 076f45c1ee [Feature] Support Lora for DMD (#755)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2025-09-08 14:18:21 -07:00
85eb7265db fix: lora_B init zeros (#781)
Co-authored-by: zbchu2 <zbchu2@iflytek.com>
Co-authored-by: Wenxuan Tan <wenxuan.tan@wisc.edu>
2025-09-05 22:56:52 -07:00
William Lin d3ceb67e66 [misc] Update Slack invite link (#786) 2025-09-05 12:16:18 -07:00
Zhang Peiyuan 7ac153a5ca Update WeChat Link 2025-09-05 11:40:47 -07:00
William Lin d1e7aa0abd [CI] Add ssim test for causal inference (#784) 2025-09-05 01:23:01 -07:00
William Lin 2d846c55a1 [misc] Improve text encoding stage (#774) 2025-09-04 17:51:27 -07:00
Jinzhe Pan b318063c0a [Preprocess][Fix] video quality issue (#773) 2025-09-03 20:47:33 -07:00
Jinzhe Pan 4aa307be55 [Preprocess][Feat] support torchvision to load video in new preprocessing (#761) 2025-09-01 23:37:01 -07:00
90 changed files with 4055 additions and 3628 deletions
+23
View File
@@ -176,3 +176,26 @@ 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,6 +109,15 @@ 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
+2 -2
View File
@@ -235,7 +235,7 @@ jobs:
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
training-test:
needs: change-filter
if: >-
@@ -373,4 +373,4 @@ jobs:
JOB_IDS: '["encoder-test", "vae-test", "transformer-test", "ssim-test-py3.10", "ssim-test-py3.11", "ssim-test-py3.12", "training-test", "training-test-VSA", "inference-test-STA", "precision-test-STA", "precision-test-VSA"]'
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
GITHUB_RUN_ID: ${{ github.run_id }}
run: python .github/scripts/runpod_cleanup.py
run: python .github/scripts/runpod_cleanup.py
-2
View File
@@ -64,5 +64,3 @@ docs/source/distillation/examples/
!docs/source/_static/images/**/*.png
!comfyui/assets/**/*.png
!comfyui/assets/**/*.gif
dmd_t2v_output/
+1 -1
View File
@@ -7,7 +7,7 @@
FastVideo features an end-to-end unified pipeline for accelerating diffusion models, starting from data preprocessing to model training, finetuning, distillation, and inference. FastVideo is designed to be modular and extensible, allowing users to easily add new optimizations and techniques. Whether it is training-free optimizations or post-training optimizations, FastVideo has you covered.
<p align="center">
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-38u6p1jqe-yDI1QJOCEnbtkLoaI5bjZQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/rG0QpZdw" target="_blank"> <b> WeChat </b> </a> |
| 🕹️ <a href="https://fastwan.fastvideo.org/"<b>Online Demo</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start.html"><b> Quick Start</b></a> | 🤗 <a href="https://huggingface.co/collections/FastVideo/fastwan-6886a305d9799c8cd1496408" target="_blank"><b>FastWan</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3csdw1isz-Euq8_Q8~baewG8hxjXs2gQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://ibb.co/S7HLCSTh" target="_blank"> <b> WeChat </b> </a> |
</p>
<div align="center">
+32
View File
@@ -0,0 +1,32 @@
# 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
@@ -0,0 +1,24 @@
# 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=[]
)
@@ -0,0 +1,97 @@
# 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
@@ -0,0 +1,2 @@
# SPDX-License-Identifier: Apache-2.0
from .vmoba import moba_attn_varlen, process_moba_input, process_moba_output
+860
View File
@@ -0,0 +1,860 @@
# 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')
@@ -1,155 +0,0 @@
#!/bin/bash
#SBATCH --job-name=t2v
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:1
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=dmd_t2v_output/t2v_%j.out
#SBATCH --error=dmd_t2v_output/t2v_%j.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate wei-fv
# Basic Info
export WANDB_MODE="online"
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export CUDA_VISIBLE_DEVICES=$SLURM_LOCALID
export TOKENIZERS_PARALLELISM=false
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# Configs
NUM_GPUS=1
# Model paths for DMD distillation:
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-14B-Diffusers" # Teacher model
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
DATA_DIR="data/crush-smol-single_processed_t2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="data/crush-smol-single_processed_t2v/validation.json"
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd
--output_dir "checkpoints/SFwan_t2v_finetune"
--max_train_steps 500
--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 81
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frame_per_block 3
--enable_gradient_masking
--gradient_mask_last_n_frames 21
)
# Parallel arguments
parallel_args=(
--num_gpus 1 # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim 1
)
# Model arguments
model_args=(
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
--generator_model_path $GENERATOR_MODEL_PATH
--real_score_model_path $REAL_SCORE_MODEL_PATH
--fake_score_model_path $FAKE_SCORE_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 10
--validation_sampling_steps "4"
--validation_guidance_scale "6.0" # not used for dmd inference
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--training_state_checkpointing_steps 50
--weight_only_checkpointing_steps 50
--weight_decay 0.01
--betas '0.0,0.999'
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.0
--dit_precision "fp32"
--flow_shift 5
--seed 1000
--use_ema True
--ema_decay 0.99
--ema_start_step 100
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
)
# Self-forcing DMD arguments
dmd_args=(
--dmd_denoising_steps '1000,750,500,250'
--min_timestep_ratio 0.02
--max_timestep_ratio 0.98
--dfake_gen_update_ratio 5
--real_score_guidance_scale 3.0
--fake_score_learning_rate 8e-6
--fake_score_betas '0.0,0.999'
)
srun torchrun \
--nnodes $SLURM_JOB_NUM_NODES \
--nproc_per_node $NUM_GPUS \
--node_rank $SLURM_PROCID \
--rdzv_backend=c10d \
--rdzv_endpoint="$MASTER_ADDR:$MASTER_PORT" \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}"
@@ -1,3 +0,0 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -1,43 +0,0 @@
#!/bin/bash
# Download the full dataset first
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
# Create a single-example dataset for debugging
SINGLE_EXAMPLE_DIR="data/crush-smol-single"
mkdir -p "$SINGLE_EXAMPLE_DIR/videos"
# Copy the specific video that matches the validation.json style (macaron crushing)
cp "data/crush-smol/videos/7P02AihYkCU-Scene-005.mp4" "$SINGLE_EXAMPLE_DIR/videos/"
# Create a single-line videos.txt
echo "videos/7P02AihYkCU-Scene-005.mp4" > "$SINGLE_EXAMPLE_DIR/videos.txt"
# Create a single-line prompt.txt with the macaron crushing prompt
echo "PIKA_CRUSH 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." > "$SINGLE_EXAMPLE_DIR/prompt.txt"
# Generate the JSON file and merge.txt for the single example
python scripts/dataset_preparation/prepare_json_file.py --data_folder "$SINGLE_EXAMPLE_DIR" --output "videos2caption.json"
# Create a validation.json that uses the same example for consistency
cat > "$SINGLE_EXAMPLE_DIR/validation.json" << 'EOF'
{
"data": [
{
"caption": "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.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 81
}
]
}
EOF
echo "Single example dataset created at $SINGLE_EXAMPLE_DIR"
echo "Contains:"
echo "- 1 video: $(cat $SINGLE_EXAMPLE_DIR/videos.txt)"
echo "- 1 prompt: $(cat $SINGLE_EXAMPLE_DIR/prompt.txt)"
echo "- Validation file created with the same example for consistency"
@@ -1,29 +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="data/crush-smol-single/merge.txt"
OUTPUT_DIR="data/crush-smol-single_processed_t2v/"
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 1 \
--flush_frequency 1 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v"
# Copy the validation.json to the output directory for consistency
cp "data/crush-smol-single/validation.json" "$OUTPUT_DIR/"
echo "Preprocessing completed. Validation file copied to $OUTPUT_DIR/"
@@ -1,31 +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": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
"image_path": null,
"video_path": null,
"num_inference_steps": 50,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -98,6 +98,7 @@ dmd_args=(
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
--master_port $MASTER_PORT \
fastvideo/training/wan_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
@@ -0,0 +1,112 @@
#!/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[@]}"
@@ -0,0 +1,47 @@
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.
@@ -0,0 +1,93 @@
#!/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[@]}"
@@ -3,14 +3,14 @@
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_TYPE="wan"
DATA_MERGE_PATH="data/crush-smol/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
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 8 \
--preprocess_video_batch_size 1 \
--seed 42 \
--max_height 480 \
--max_width 832 \
@@ -21,4 +21,4 @@ torchrun --nproc_per_node=$GPU_NUM \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v"
--preprocess_task "ode_trajectory"
@@ -0,0 +1,40 @@
{
"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
}
]
}
@@ -10,6 +10,7 @@ torchrun --nproc_per_node=$GPU_NUM \
--model_path $MODEL_PATH \
--mode preprocess \
--workload_type t2v \
--preprocess.video_loader_type torchvision \
--preprocess.dataset_type merged \
--preprocess.dataset_path $DATASET_PATH \
--preprocess.dataset_output_dir $OUTPUT_DIR \
@@ -28,4 +28,4 @@
"num_frames": 77
}
]
}
}
+214
View File
@@ -0,0 +1,214 @@
# 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
@@ -0,0 +1,16 @@
{
"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
}
@@ -0,0 +1,16 @@
{
"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
}
+34
View File
@@ -32,6 +32,29 @@ class DatasetType(str, Enum):
return [dataset_type.value for dataset_type in cls]
class VideoLoaderType(str, Enum):
"""
Enumeration for different video loaders.
"""
TORCHCODEC = "torchcodec"
TORCHVISION = "torchvision"
@classmethod
def from_string(cls, value: str) -> "VideoLoaderType":
"""Convert string to VideoLoader enum."""
try:
return cls(value.lower())
except ValueError:
raise ValueError(
f"Invalid video loader: {value}. Must be one of: {', '.join([m.value for m in cls])}"
) from None
@classmethod
def choices(cls) -> list[str]:
"""Get all available choices as strings for argparse."""
return [video_loader.value for video_loader in cls]
@dataclasses.dataclass
class PreprocessConfig:
"""Configuration for preprocessing operations."""
@@ -51,6 +74,7 @@ class PreprocessConfig:
flush_frequency: int = 256
# Video processing parameters
video_loader_type: VideoLoaderType = VideoLoaderType.TORCHCODEC
max_height: int = 480
max_width: int = 848
num_frames: int = 163
@@ -120,6 +144,12 @@ class PreprocessConfig:
help="How often to save to parquet files")
# Video processing parameters
preprocess_args.add_argument(
f"--{prefix_with_dot}video-loader-type",
type=str,
choices=VideoLoaderType.choices(),
default=PreprocessConfig.video_loader_type.value,
help="Type of the video loader")
preprocess_args.add_argument(f"--{prefix_with_dot}max-height",
type=int,
default=PreprocessConfig.max_height,
@@ -174,6 +204,10 @@ class PreprocessConfig:
if 'dataset_type' in kwargs and isinstance(kwargs['dataset_type'], str):
kwargs['dataset_type'] = DatasetType.from_string(
kwargs['dataset_type'])
if 'video_loader_type' in kwargs and isinstance(
kwargs['video_loader_type'], str):
kwargs['video_loader_type'] = VideoLoaderType.from_string(
kwargs['video_loader_type'])
preprocess_config = cls()
if not update_config_from_args(
+7 -3
View File
@@ -15,9 +15,13 @@ 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.SLIDING_TILE_ATTN,
AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.VMOBA_ATTN,
)
hidden_size: int = 0
num_attention_heads: int = 0
+1 -1
View File
@@ -109,4 +109,4 @@ class WanVideoArchConfig(DiTArchConfig):
class WanVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=WanVideoArchConfig)
prefix: str = "Wan"
prefix: str = "Wan"
+3 -2
View File
@@ -4,12 +4,13 @@ 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 (WanI2V480PConfig, WanI2V720PConfig,
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
WanI2V480PConfig, WanI2V720PConfig,
WanT2V480PConfig, WanT2V720PConfig)
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
"get_pipeline_config_cls_from_name"
"SelfForcingWanT2V480PConfig", "get_pipeline_config_cls_from_name"
]
+12 -8
View File
@@ -7,12 +7,14 @@ from collections.abc import Callable
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.wan import (FastWan2_1_T2V_480P_Config,
FastWan2_2_TI2V_5B_Config,
SelfForcingWanT2V480PConfig,
WanI2V480PConfig, WanI2V720PConfig,
WanT2V480PConfig, WanT2V720PConfig)
# 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)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
verify_model_config_and_directory)
@@ -33,10 +35,10 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": FastWan2_2_TI2V_5B_Config,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VConfig,
"FastVideo/Wan2.1-VSA-T2V-14B-720P-Diffusers": WanT2V720PConfig,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": WanT2V720PConfig,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": WanT2V480PConfig,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": WanI2V480PConfig,
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers": SelfForcingWanT2V480PConfig,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_Config,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
# Add other specific weight variants
}
@@ -46,6 +48,7 @@ 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
}
@@ -58,6 +61,7 @@ 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
}
+13 -12
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: int = 3
flow_shift: float | None = 3.0
# 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: int = 5
flow_shift: float | None = 5.0
@dataclass
@@ -94,7 +94,7 @@ class WanI2V720PConfig(WanI2V480PConfig):
# WanConfig-specific parameters with defaults
# Denoising stage
flow_shift: int = 5
flow_shift: float | None = 5.0
@dataclass
@@ -104,7 +104,7 @@ class FastWan2_1_T2V_480P_Config(WanT2V480PConfig):
# WanConfig-specific parameters with defaults
# Denoising stage
flow_shift: int = 8
flow_shift: float | None = 8.0
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: int = 5
flow_shift: float | None = 5.0
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: int = 5
flow_shift: float | None = 5.0
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 757, 522])
@@ -146,6 +146,7 @@ class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
@dataclass
class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
is_causal: bool = True
flow_shift: int = 5
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,6 +47,8 @@ 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"
@@ -191,6 +193,25 @@ 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
View File
@@ -37,6 +37,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
# Wan2.2
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
+8 -2
View File
@@ -4,7 +4,7 @@ from torchvision.transforms import Lambda
from fastvideo.dataset.parquet_dataset_map_style import (
build_parquet_map_style_dataloader)
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset, TextDataset
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
TemporalRandomCrop)
from fastvideo.dataset.validation_dataset import ValidationDataset
@@ -37,9 +37,15 @@ 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,
args=args,
seed=args.seed)
__all__ = [
"build_parquet_map_style_dataloader", "ValidationDataset",
"VideoCaptionMergedDataset"
"VideoCaptionMergedDataset", "TextDataset"
]
+264
View File
@@ -0,0 +1,264 @@
"""
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
+24
View File
@@ -50,6 +50,7 @@ pyarrow_schema_i2v = pa.schema([
pa.field("fps", pa.float64()),
])
pyarrow_schema_t2v = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
@@ -78,3 +79,26 @@ pyarrow_schema_t2v = pa.schema([
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
pyarrow_schema_ode_trajectory_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()),
# --- 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()), # Always 'text' for text-only
])
+131
View File
@@ -628,3 +628,134 @@ class VideoCaptionMergedDataset(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"]
class TextDataset(torch.utils.data.IterableDataset,
torch.distributed.checkpoint.stateful.Stateful):
"""
Text-only dataset for processing prompts from a simple text file.
Assumes that data_merge_path is a text file with one prompt per line:
A cat playing with a ball
A dog running in the park
A person cooking dinner
...
This dataset processes text data through text encoding stages only.
"""
def __init__(self,
data_merge_path: str,
args,
start_idx: int = 0,
seed: int = 42):
self.data_merge_path = data_merge_path
self.start_idx = start_idx
self.args = args
self.seed = seed
# Initialize tokenizer
tokenizer_path = os.path.join(args.model_path, "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
# Initialize text encoding stage
self.text_encoding_stage = TextEncodingStage(
tokenizer=tokenizer,
text_max_length=args.text_max_length,
cfg_rate=getattr(args, 'training_cfg_rate', 0.0),
seed=self.seed)
# Process text data
self.processed_batches = self._process_text_data()
def _load_text_data(self) -> list[str]:
"""Load text prompts from file."""
prompts = []
with open(self.data_merge_path, 'r', encoding='utf-8') as f:
for line in f:
line = line.strip()
if line: # Skip empty lines
prompts.append(line)
logger.info(f"Loaded {len(prompts)} text prompts from {self.data_merge_path}")
return prompts
def _process_text_data(self) -> list[PreprocessBatch]:
"""Process the text prompts through text encoding stage."""
raw_prompts = self._load_text_data()
processed_batches = []
for idx, prompt in enumerate(raw_prompts):
# Create a text-only batch with dummy path
batch = PreprocessBatch(
path=f"text_prompt_{idx}",
cap=[prompt], # TextEncodingStage expects a list
resolution=None,
fps=None,
duration=None,
num_frames=0,
sample_frame_index=None,
sample_num_frames=0
)
processed_batches.append(batch)
logger.info(f"Processed {len(processed_batches)} text batches")
return processed_batches
def __iter__(self):
"""Iterator for the dataset."""
# Set up distributed sampling if needed
if torch.distributed.is_available() and torch.distributed.is_initialized():
rank = torch.distributed.get_rank()
world_size = torch.distributed.get_world_size()
else:
rank = 0
world_size = 1
# Calculate chunk for this rank
total_items = len(self.processed_batches)
items_per_rank = math.ceil(total_items / world_size)
start_idx = rank * items_per_rank + self.start_idx
end_idx = min(start_idx + items_per_rank, total_items)
# Yield items for this rank
for idx in range(start_idx, end_idx):
if idx < len(self.processed_batches):
yield self._get_item(idx)
def _get_item(self, idx: int) -> dict:
"""Get a single processed text item."""
batch = self.processed_batches[idx]
# Apply text encoding stage
batch = self.text_encoding_stage.process(batch)
# Build result dictionary for text-only processing with required schema fields
result = {
"text": batch.text,
"input_ids": batch.input_ids,
"cond_mask": batch.cond_mask,
"path": batch.path,
# Required schema fields for ODE trajectory processing
"id": f"text_{idx}",
"file_name": batch.path,
"caption": batch.text,
"media_type": "text",
"width": 1,
"height": 1,
"num_frames": 0,
"duration_sec": 0.0,
"fps": 0.0,
}
return result
def state_dict(self) -> dict[str, Any]:
"""Return state dict for checkpointing."""
return {"processed_batches": self.processed_batches}
def load_state_dict(self, state_dict: dict[str, Any]) -> None:
"""Load state dict from checkpoint."""
self.processed_batches = state_dict["processed_batches"]
+1 -1
View File
@@ -5,7 +5,7 @@ import numpy as np
import torch
def pad(t: torch.Tensor, padding_length: int) -> torch.Tensor:
def pad(t: torch.Tensor, padding_length: int) -> tuple[torch.Tensor, torch.Tensor]:
"""
Pad or crop an embedding [L, D] to exactly padding_length tokens.
Return:
+3
View File
@@ -344,6 +344,9 @@ 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,
+25 -83
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
@@ -119,7 +119,7 @@ class FastVideoArgs:
output_type: str = "pil"
# CPU offload parameters
dit_cpu_offload: bool = False
dit_cpu_offload: bool = True
use_fsdp_inference: bool = True
text_encoder_cpu_offload: bool = True
image_encoder_cpu_offload: bool = True
@@ -139,6 +139,10 @@ 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
@@ -166,6 +170,16 @@ 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
@@ -591,11 +605,6 @@ class TrainingArgs(FastVideoArgs):
pretrained_model_name_or_path: str = ""
dit_model_name_or_path: str = ""
# DMD model paths - separate paths for each network
generator_model_path: str = "" # path for generator (student) model
real_score_model_path: str = "" # path for real score (teacher) model
fake_score_model_path: str = "" # path for fake score (critic) model
# diffusion setting
ema_decay: float = 0.0
ema_start_step: int = 0
@@ -618,7 +627,6 @@ class TrainingArgs(FastVideoArgs):
checkpoints_total_limit: int = 0
checkpointing_steps: int = 0
resume_from_checkpoint: str = "" # specify the checkpoint folder to resume from
init_weights_from_safetensors: str = "" # path to safetensors file for initial weight loading
# optimizer & scheduler
num_train_epochs: int = 0
@@ -650,7 +658,6 @@ class TrainingArgs(FastVideoArgs):
linear_quadratic_threshold: float = 0.0
linear_range: float = 0.0
weight_decay: float = 0.0
betas: str = "0.9,0.999" # betas for optimizer, format: "beta1,beta2"
use_ema: bool = False
multi_phased_distill_schedule: str = ""
pred_decay_weight: float = 0.0
@@ -671,26 +678,16 @@ class TrainingArgs(FastVideoArgs):
# distillation args
generator_update_interval: int = 5
dfake_gen_update_ratio: int = 5 # self-forcing: how often to train generator vs critic
min_timestep_ratio: float = 0.2
max_timestep_ratio: float = 0.98
real_score_guidance_scale: float = 3.5
fake_score_learning_rate: float = 0.0 # separate learning rate for fake_score_transformer, if 0.0, use learning_rate
fake_score_lr_scheduler: str = "constant" # separate lr scheduler for fake_score_transformer, if not set, use lr_scheduler
fake_score_betas: str = "0.9,0.999" # betas for fake score optimizer, format: "beta1,beta2"
training_state_checkpointing_steps: int = 0 # for resuming training
weight_only_checkpointing_steps: int = 0 # for inference
log_visualization: bool = False
# simulate generator forward to match inference
simulate_generator_forward: bool = False
warp_denoising_step: bool = False
# Self-forcing specific arguments
num_frame_per_block: int = 3
independent_first_frame: bool = False
enable_gradient_masking: bool = True
gradient_mask_last_n_frames: int = 21
validate_cache_structure: bool = False # Debug flag for cache validation
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
@@ -792,20 +789,6 @@ class TrainingArgs(FastVideoArgs):
type=str,
help="Directory to cache models")
# DMD model paths - separate paths for each network
parser.add_argument(
"--generator-model-path",
type=str,
help="Path to generator (student) model for DMD distillation")
parser.add_argument(
"--real-score-model-path",
type=str,
help="Path to real score (teacher) model for DMD distillation")
parser.add_argument(
"--fake-score-model-path",
type=str,
help="Path to fake score (critic) model for DMD distillation")
# Diffusion settings
parser.add_argument("--ema-decay",
type=float,
@@ -876,10 +859,6 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--resume-from-checkpoint",
type=str,
help="Path to checkpoint to resume from")
parser.add_argument(
"--init-weights-from-safetensors",
type=str,
help="Path to safetensors file for initial weight loading")
parser.add_argument("--logging-dir",
type=str,
help="Directory for logging")
@@ -984,10 +963,6 @@ class TrainingArgs(FastVideoArgs):
help="Linear quadratic threshold")
parser.add_argument("--linear-range", type=float, help="Linear range")
parser.add_argument("--weight-decay", type=float, help="Weight decay")
parser.add_argument("--betas",
type=str,
default=TrainingArgs.betas,
help="Betas for optimizer (format: 'beta1,beta2')")
parser.add_argument("--use-ema",
action=StoreBoolean,
help="Whether to use EMA")
@@ -1024,18 +999,20 @@ 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,
default=TrainingArgs.generator_update_interval,
help="Ratio of student updates to critic updates.")
parser.add_argument(
"--dfake-gen-update-ratio",
type=int,
default=TrainingArgs.dfake_gen_update_ratio,
help=
"Self-forcing: How often to train generator vs critic (train generator every N steps)."
)
parser.add_argument("--min-timestep-ratio",
type=float,
default=TrainingArgs.min_timestep_ratio,
@@ -1052,11 +1029,6 @@ class TrainingArgs(FastVideoArgs):
type=float,
default=TrainingArgs.fake_score_learning_rate,
help="Learning rate for fake score transformer")
parser.add_argument(
"--fake-score-betas",
type=str,
default=TrainingArgs.fake_score_betas,
help="Betas for fake score optimizer (format: 'beta1,beta2')")
parser.add_argument(
"--fake-score-lr-scheduler",
type=str,
@@ -1069,36 +1041,6 @@ class TrainingArgs(FastVideoArgs):
"--simulate-generator-forward",
action=StoreBoolean,
help="Whether to simulate generator forward to match inference")
parser.add_argument(
"--warp-denoising-step",
action=StoreBoolean,
help=
"Whether to warp denoising step according to the scheduler time shift"
)
# Self-forcing specific arguments
parser.add_argument(
"--num-frame-per-block",
type=int,
default=TrainingArgs.num_frame_per_block,
help="Number of frames per block for causal generation")
parser.add_argument(
"--independent-first-frame",
action=StoreBoolean,
help="Whether the first frame is independent in causal generation")
parser.add_argument(
"--enable-gradient-masking",
action=StoreBoolean,
help="Whether to enable frame-level gradient masking")
parser.add_argument(
"--gradient-mask-last-n-frames",
type=int,
default=TrainingArgs.gradient_mask_last_n_frames,
help="Number of last frames to enable gradients for")
parser.add_argument(
"--validate-cache-structure",
action=StoreBoolean,
help="Whether to validate KV cache structure (debug flag)")
return parser
+60 -6
View File
@@ -100,7 +100,16 @@ class ScaleResidual(nn.Module):
def forward(self, residual: torch.Tensor, x: torch.Tensor,
gate: torch.Tensor) -> torch.Tensor:
"""Apply gated residual connection."""
return residual + x * gate
# 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
# adapted from Diffusers: https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/normalization.py
@@ -159,7 +168,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, shift: torch.Tensor,
gate: torch.Tensor | int, shift: torch.Tensor,
scale: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""
Apply gated residual connection, followed by layernorm and
@@ -171,12 +180,41 @@ 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
residual_output = residual + x * gate
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]
# Apply normalization
normalized = self.norm(residual_output)
# Apply scale and shift
modulated = normalized * (1.0 + scale) + 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
return modulated, residual_output
@@ -218,8 +256,24 @@ 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:
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
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)
else:
return normalized * (1.0 + scale) + shift
# 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
+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.kaiming_uniform_(self.lora_B, a=math.sqrt(5))
torch.nn.init.zeros_(self.lora_B)
else:
self.lora_A = None
self.lora_B = None
+1
View File
@@ -33,6 +33,7 @@ from fastvideo.logger import init_logger
logger = init_logger(__name__)
def _rotate_neox(x: torch.Tensor) -> torch.Tensor:
x1 = x[..., :x.shape[-1] // 2]
x2 = x[..., x.shape[-1] // 2:]
+20 -11
View File
@@ -244,19 +244,26 @@ 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=1)
6, dim=2)
# *_msa.shape: [batch_size, num_frames, 1, inner_dim]
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
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)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
@@ -487,8 +494,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, encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
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)
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
@@ -526,8 +533,9 @@ class CausalWanTransformer3DModel(BaseDiT):
**causal_kwargs)
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2,
dim=2)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
@@ -596,8 +604,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, encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
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)
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
@@ -623,8 +631,9 @@ class CausalWanTransformer3DModel(BaseDiT):
block_mask=self.block_mask)
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2,
dim=2)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
@@ -430,16 +430,6 @@ class TransformerLoader(ComponentLoader):
if not safetensors_list:
raise ValueError(f"No safetensors files found in {model_path}")
# Check if we should use custom initialization weights
custom_weights_path = getattr(fastvideo_args, 'init_weights_from_safetensors', None)
use_custom_weights = (custom_weights_path and os.path.exists(custom_weights_path) and
fastvideo_args.training_mode and
not hasattr(fastvideo_args, '_loading_teacher_critic_model'))
if use_custom_weights:
logger.info("Using custom initialization weights from: %s", custom_weights_path)
safetensors_list = [custom_weights_path]
logger.info("Loading model from %s safetensors files in %s",
len(safetensors_list), model_path)
-7
View File
@@ -248,13 +248,6 @@ def load_model_from_full_model_state_dict(
sharded_sd = {}
custom_param_sd, reverse_param_names_mapping = hf_to_custom_state_dict(
full_sd_iterator, param_names_mapping) # type: ignore
print(custom_param_sd.keys())
print("--------------------------------")
print("--------------------------------")
print("--------------------------------")
print("--------------------------------")
print("--------------------------------")
print(meta_sd.keys())
for target_param_name, full_tensor in custom_param_sd.items():
meta_sharded_param = meta_sd.get(target_param_name)
if meta_sharded_param is None:
+3
View File
@@ -61,6 +61,9 @@ _SCHEDULERS = {
"FlowMatchEulerDiscreteScheduler"),
"UniPCMultistepScheduler":
("schedulers", "scheduling_unipc_multistep", "UniPCMultistepScheduler"),
"SelfForcingFlowMatchScheduler":
("schedulers", "scheduling_self_forcing_flow_match",
"SelfForcingFlowMatchScheduler"),
}
_FAST_VIDEO_MODELS = {
@@ -0,0 +1,122 @@
# 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
+23 -1
View File
@@ -145,8 +145,30 @@ 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]
"""
timestep = timestep.expand(noise_input_latent.shape[0])
# 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]
dtype = pred_noise.dtype
device = pred_noise.device
pred_noise = pred_noise.float().to(device)
@@ -16,8 +16,7 @@ from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
CausalDMDDenosingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
TextEncodingStage)
# isort: on
logger = init_logger(__name__)
@@ -48,10 +47,6 @@ 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, torch.nn.Module] = {}
modules: dict[str, Any] = {}
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 name == "transformer":
if "transformer" in name:
module.requires_grad_(True)
else:
module.requires_grad_(False)
@@ -136,11 +136,9 @@ class ComposedPipelineBase(ABC):
kwargs['model_path'] = model_path
fastvideo_args = FastVideoArgs.from_kwargs(**kwargs)
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
else:
assert args is not None, "args must be provided for training mode"
fastvideo_args = TrainingArgs.from_cli_args(args)
logger.info("training args in from_pretrained: %s", fastvideo_args)
# TODO(will): fix this so that its not so ugly
fastvideo_args.model_path = model_path
for key, value in kwargs.items():
+36 -12
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 init_device_mesh
from torch.distributed.device_mesh import DeviceMesh, init_device_mesh
from torch.distributed.tensor import DTensor
from fastvideo.distributed import get_local_torch_device
@@ -32,6 +32,7 @@ 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()
@@ -81,6 +82,17 @@ 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:
@@ -88,18 +100,12 @@ 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"])
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))
set_lora_grads(self.lora_layers, device_mesh)
set_lora_grads(self.lora_layers_critic, device_mesh)
def convert_to_lora_layers(self) -> None:
"""
@@ -131,6 +137,24 @@ 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
@@ -224,4 +248,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()
+12 -3
View File
@@ -147,7 +147,12 @@ class ForwardBatch:
modules: dict[str, Any] = field(default_factory=dict)
# Final output (after pipeline completion)
output: Any = None
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
# Extra parameters that might be needed by specific pipeline implementations
extra: dict[str, Any] = field(default_factory=dict)
@@ -206,6 +211,10 @@ 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
@@ -241,5 +250,5 @@ class TrainingBatch:
@dataclass
class PreprocessBatch(ForwardBatch):
video_loader: list["VideoDecoder"] = field(default_factory=list)
video_file_name: list[str] = field(default_factory=list)
video_loader: list["VideoDecoder"] | list[str] = field(default_factory=list)
video_file_name: list[str] = field(default_factory=list)
@@ -1,7 +1,5 @@
# SPDX-License-Identifier: Apache-2.0
import multiprocessing
import os
from concurrent.futures import ProcessPoolExecutor
from typing import Any
import numpy as np
@@ -19,6 +17,8 @@ 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,10 +54,14 @@ class BasePreprocessPipeline(ComposedPipelineBase):
"""Get additional features specific to the pipeline type. Override in subclasses."""
return {}
def get_schema_fields(self) -> list[str]:
"""Get the schema fields for the pipeline type. Override in subclasses."""
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()]
def create_record_for_schema(self,
preprocess_batch: PreprocessBatch,
schema: pa.Schema,
@@ -400,166 +404,26 @@ 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")
# 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())
table = records_to_table(batch_data, self.get_pyarrow_schema())
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)
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)
logger.info("Collected batch with %s samples", len(table))
if num_processed_samples >= args.flush_frequency:
self._flush_tables(num_processed_samples, args,
combined_parquet_dir)
written = self.dataset_writer.flush()
logger.info("Flushed %s samples to parquet", written)
num_processed_samples = 0
self.all_tables = []
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
def _final_flush_if_any(self):
if hasattr(self, 'dataset_writer'):
self.dataset_writer.flush()
@@ -40,9 +40,9 @@ class PreprocessPipeline_I2V(BasePreprocessPipeline):
image_processor=self.get_module("image_processor"),
))
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_pyarrow_schema(self):
"""Return the PyArrow schema for I2V pipeline."""
return pyarrow_schema_i2v
def get_extra_features(self, valid_data: dict[str, Any],
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
@@ -0,0 +1,399 @@
# SPDX-License-Identifier: Apache-2.0
"""
ODE Trajectory Data Preprocessing pipeline implementation.
This module contains an implementation of the ODE Trajectory Data Preprocessing pipeline
using the modular pipeline architecture.
Sec 4.3 of CausVid paper: https://arxiv.org/pdf/2412.07772
"""
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.configs.sample import SamplingParam
from fastvideo.dataset import gettextdataset
from fastvideo.dataset.dataloader.schema import (
pyarrow_schema_ode_trajectory_text_only)
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.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
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)
logger = init_logger(__name__)
class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"""ODE Trajectory preprocessing pipeline implementation."""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
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 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",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
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."""
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],
}
# 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
]
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]
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,
strict=False)):
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.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.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.do_classifier_free_guidance = True
result_batch = self.input_validation_stage(
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.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
extra_features = {
"trajectory_latents": trajectory_latents,
"trajectory_timesteps": trajectory_timesteps
}
if batch.return_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)
# 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:
video_name = os.path.basename(video_path).split(".")[0]
# Convert tensors to numpy arrays
text_embedding = prompt_embeds[idx].cpu().numpy()
# Get extra features for this sample
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
if isinstance(value, torch.Tensor):
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()
else:
sample_extra_features[key] = value[idx]
# Create record for Parquet dataset (without VAE latents for text-only)
record = self.create_text_only_record(
args,
video_name=video_name,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=sample_extra_features)
batch_data.append(record)
if batch_data:
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
table = records_to_table(batch_data,
self.get_pyarrow_schema())
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)
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.num_processed_samples = 0
# 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)
def create_text_only_record(
self,
args,
video_name: str,
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 text-only preprocessing using text-only schema."""
# Create base record using only fields from text-only schema
record = {
"id": f"text_{video_name}_{idx}",
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"file_name": video_name,
"caption": valid_data["text"][idx],
"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"]
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"][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),
})
else:
record.update({
"trajectory_timesteps_bytes": b"",
"trajectory_timesteps_shape": [],
"trajectory_timesteps_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 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 video preprocessing
self.pbar = tqdm(self.preprocess_loader_iter,
desc="Processing videos",
unit="batch",
disable=self.local_rank != 0)
# 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_text_and_trajectory(fastvideo_args, args)
EntryClass = PreprocessPipeline_ODE_Trajectory
@@ -15,9 +15,9 @@ class PreprocessPipeline_T2V(BasePreprocessPipeline):
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
def get_schema_fields(self):
"""Get the schema fields for T2V pipeline."""
return [f.name for f in pyarrow_schema_t2v]
def get_pyarrow_schema(self):
"""Return the PyArrow schema for T2V pipeline."""
return pyarrow_schema_t2v
EntryClass = PreprocessPipeline_T2V
@@ -4,9 +4,11 @@ from typing import cast
import numpy as np
import torch
import torchvision
from einops import rearrange
from torchvision import transforms
from fastvideo.configs.configs import VideoLoaderType
from fastvideo.dataset.transform import (CenterCropResizeVideo,
TemporalRandomCrop)
from fastvideo.fastvideo_args import FastVideoArgs, WorkloadType
@@ -61,7 +63,16 @@ class VideoTransformStage(PipelineStage):
else:
frame_indices = frame_indices[:self.num_frames]
video = batch.video_loader[i].get_frames_at(frame_indices).data
if fastvideo_args.preprocess_config.video_loader_type == VideoLoaderType.TORCHCODEC:
video = batch.video_loader[i].get_frames_at(frame_indices).data
elif fastvideo_args.preprocess_config.video_loader_type == VideoLoaderType.TORCHVISION:
video, _, _ = torchvision.io.read_video(batch.video_loader[i],
output_format="TCHW")
video = video[frame_indices]
else:
raise ValueError(
f"Invalid video loader type: {fastvideo_args.preprocess_config.video_loader_type}"
)
video = self.video_transform(video)
video_pixel_batch.append(video)
@@ -9,6 +9,8 @@ from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.preprocess.preprocess_pipeline_i2v import (
PreprocessPipeline_I2V)
from fastvideo.pipelines.preprocess.preprocess_pipeline_ode_trajectory import (
PreprocessPipeline_ODE_Trajectory)
from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
PreprocessPipeline_T2V)
from fastvideo.utils import maybe_download_model
@@ -24,7 +26,8 @@ def main(args) -> None:
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
"flow_shift": 5,
}
pipeline_config.update_config_from_dict(kwargs)
fastvideo_args = FastVideoArgs(
@@ -35,7 +38,17 @@ def main(args) -> None:
text_encoder_cpu_offload=False,
pipeline_config=pipeline_config,
)
PreprocessPipeline = PreprocessPipeline_I2V if args.preprocess_task == "i2v" else PreprocessPipeline_T2V
if args.preprocess_task == "t2v":
PreprocessPipeline = PreprocessPipeline_T2V
elif args.preprocess_task == "i2v":
PreprocessPipeline = PreprocessPipeline_I2V
elif args.preprocess_task == "ode_trajectory":
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
else:
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}")
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)
+43 -7
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,7 +36,6 @@ 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
@@ -71,8 +70,16 @@ class CausalDMDDenosingStage(DenoisingStage):
# Timesteps for DMD
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long,
device=get_local_torch_device())
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)
# Image kwargs (kept empty unless caller provides compatible args)
image_kwargs: dict = {}
@@ -262,10 +269,14 @@ 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_expand,
t_expanded_noise,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=(pos_start_base + start_index) *
@@ -326,10 +337,11 @@ 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_context,
t_expanded_context,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=(pos_start_base + start_index) *
@@ -407,3 +419,27 @@ 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
+91 -48
View File
@@ -50,6 +50,63 @@ 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,
@@ -59,13 +116,28 @@ 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 information.
fastvideo_args: The inference arguments.
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
Returns:
The batch with decoded outputs.
Modified batch with:
- output: Decoded frames (batch, channels, frames, height, width) as CPU float32
- trajectory_decoded (if requested): List of decoded frames per timestep
"""
# 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()
@@ -75,58 +147,29 @@ class DecodingStage(PipelineStage):
pipeline.add_module("vae", self.vae)
fastvideo_args.model_loaded["vae"] = True
self.vae = self.vae.to(get_local_torch_device())
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
frames = batch.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
frames = self.decode(batch.latents, fastvideo_args)
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)
# 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())
# Convert to CPU float32 for compatibility
image = image.cpu().float()
frames = frames.cpu().float()
# Update batch with decoded image
batch.output = image
batch.output = frames
# Offload models if needed
if hasattr(self, 'maybe_free_model_hooks'):
+74 -7
View File
@@ -40,6 +40,13 @@ 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)
@@ -77,6 +84,7 @@ class DenoisingStage(PipelineStage):
supported_attention_backends=(
AttentionBackendEnum.SLIDING_TILE_ATTN,
AttentionBackendEnum.VIDEO_SPARSE_ATTN,
AttentionBackendEnum.VMOBA_ATTN,
AttentionBackendEnum.FLASH_ATTN, AttentionBackendEnum.TORCH_SDPA
) # hack
)
@@ -149,7 +157,8 @@ class DenoisingStage(PipelineStage):
# Prepare image latents and embeddings for I2V generation
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert torch.isnan(image_embeds[0]).sum() == 0
assert not torch.isnan(
image_embeds[0]).any(), "image_embeds contains nan"
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
@@ -186,11 +195,13 @@ class DenoisingStage(PipelineStage):
# Get latents and embeddings
latents = batch.latents
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
assert not torch.isnan(
prompt_embeds[0]).any(), "prompt_embeds contains nan"
if batch.do_classifier_free_guidance:
neg_prompt_embeds = batch.negative_prompt_embeds
assert neg_prompt_embeds is not None
assert torch.isnan(neg_prompt_embeds[0]).sum() == 0
assert not torch.isnan(
neg_prompt_embeds[0]).any(), "neg_prompt_embeds contains nan"
# (Wan2.2) Calculate timestep to switch from high noise expert to low noise expert
if fastvideo_args.boundary_ratio is not None:
@@ -236,6 +247,10 @@ 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):
@@ -268,6 +283,9 @@ 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()
@@ -280,7 +298,6 @@ 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)
@@ -325,6 +342,31 @@ 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
@@ -389,6 +431,11 @@ 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
@@ -396,9 +443,28 @@ 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
@@ -737,7 +803,8 @@ class DmdDenoisingStage(DenoisingStage):
video_raw_latent_shape = latents.shape
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
assert not torch.isnan(
prompt_embeds[0]).any(), "prompt_embeds contains nan"
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long,
@@ -776,7 +843,8 @@ class DmdDenoisingStage(DenoisingStage):
batch.image_latent.permute(0, 2, 1, 3, 4)
],
dim=2).to(target_dtype)
assert torch.isnan(latent_model_input).sum() == 0
assert not torch.isnan(
latent_model_input).any(), "latent_model_input contains nan"
# Prepare inputs for transformer
t_expand = t.repeat(latent_model_input.shape[0])
@@ -838,7 +906,6 @@ class DmdDenoisingStage(DenoisingStage):
**pos_cond_kwargs,
).permute(0, 2, 1, 3, 4)
from fastvideo.models.utils import pred_noise_to_pred_video
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
-18
View File
@@ -92,24 +92,6 @@ class EncodingStage(PipelineStage):
latents = latents.to(vae_dtype)
latents = self.vae.encode(latents).mean
# Apply shifting if needed (reverse of decoding)
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
# Apply scaling factor
if (hasattr(self.vae, "scaling_factor")
and self.vae.scaling_factor is not None):
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
# Update batch with encoded latents
batch.latents = latents
+195 -51
View File
@@ -59,58 +59,37 @@ class TextEncodingStage(PipelineStage):
assert len(self.text_encoders) == len(
fastvideo_args.pipeline_config.text_encoder_configs)
for tokenizer, text_encoder, encoder_config, preprocess_func, postprocess_func in zip(
self.tokenizers,
self.text_encoders,
fastvideo_args.pipeline_config.text_encoder_configs,
fastvideo_args.pipeline_config.preprocess_text_funcs,
fastvideo_args.pipeline_config.postprocess_text_funcs,
strict=True):
# Encode positive prompt with all available encoders
assert batch.prompt is not None
prompt_text: str | list[str] = batch.prompt
all_indices: list[int] = list(range(len(self.text_encoders)))
prompt_embeds_list, prompt_masks_list = self.encode_text(
prompt_text,
fastvideo_args,
encoder_index=all_indices,
return_attention_mask=True,
)
for pe in prompt_embeds_list:
batch.prompt_embeds.append(pe)
if batch.prompt_attention_mask is not None:
for am in prompt_masks_list:
batch.prompt_attention_mask.append(am)
assert isinstance(batch.prompt, str | list)
if isinstance(batch.prompt, str):
batch.prompt = [batch.prompt]
texts = []
for prompt_str in batch.prompt:
texts.append(preprocess_func(prompt_str))
text_inputs = tokenizer(texts,
**encoder_config.tokenizer_kwargs).to(
get_local_torch_device())
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs = text_encoder(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
)
prompt_embeds = postprocess_func(outputs)
batch.prompt_embeds.append(prompt_embeds)
if batch.prompt_attention_mask is not None:
batch.prompt_attention_mask.append(attention_mask)
if batch.do_classifier_free_guidance:
assert isinstance(batch.negative_prompt, str)
negative_text = preprocess_func(batch.negative_prompt)
negative_text_inputs = tokenizer(
negative_text, **encoder_config.tokenizer_kwargs).to(
get_local_torch_device())
negative_input_ids = negative_text_inputs["input_ids"]
negative_attention_mask = negative_text_inputs["attention_mask"]
with set_forward_context(current_timestep=0,
attn_metadata=None):
negative_outputs = text_encoder(
input_ids=negative_input_ids,
attention_mask=negative_attention_mask,
output_hidden_states=True,
)
negative_prompt_embeds = postprocess_func(negative_outputs)
assert batch.negative_prompt_embeds is not None
batch.negative_prompt_embeds.append(negative_prompt_embeds)
if batch.negative_attention_mask is not None:
batch.negative_attention_mask.append(
negative_attention_mask)
# Encode negative prompt if CFG is enabled
if batch.do_classifier_free_guidance:
assert isinstance(batch.negative_prompt, str)
neg_embeds_list, neg_masks_list = self.encode_text(
batch.negative_prompt,
fastvideo_args,
encoder_index=all_indices,
return_attention_mask=True,
)
assert batch.negative_prompt_embeds is not None
for ne in neg_embeds_list:
batch.negative_prompt_embeds.append(ne)
if batch.negative_attention_mask is not None:
for nm in neg_masks_list:
batch.negative_attention_mask.append(nm)
return batch
@@ -129,6 +108,171 @@ class TextEncodingStage(PipelineStage):
V.none_or_list)
return result
@torch.no_grad()
def encode_text(
self,
text: str | list[str],
fastvideo_args: FastVideoArgs,
encoder_index: int | list[int] | None = None,
return_attention_mask: bool = False,
return_type: str = "list", # one of: "list", "dict", "stack"
device: torch.device | str | None = None,
dtype: torch.dtype | None = None,
max_length: int | None = None,
truncation: bool | None = None,
padding: bool | str | None = None,
):
"""
Encode plain text using selected text encoder(s) and return embeddings.
Args:
text: A single string or a list of strings to encode.
fastvideo_args: The inference arguments providing pipeline config,
including tokenizer and encoder settings, preprocess and postprocess
functions.
encoder_index: Encoder selector by index. Accepts an int or list of ints.
return_attention_mask: If True, also return attention masks for each
selected encoder.
return_type: "list" (default) returns a list aligned with selection;
"dict" returns a dict keyed by encoder index as a string; "stack" stacks along a
new first dimension (requires matching shapes).
device: Optional device override for inputs; defaults to local torch device.
dtype: Optional dtype to cast returned embeddings to.
max_length: Optional per-call tokenizer override.
truncation: Optional per-call tokenizer override.
padding: Optional per-call tokenizer override.
Returns:
Depending on return_type and return_attention_mask:
- list: List[Tensor] or (List[Tensor], List[Tensor])
- dict: Dict[str, Tensor] or (Dict[str, Tensor], Dict[str, Tensor])
- stack: Tensor of shape [num_encoders, ...] or a tuple with stacked
attention masks
"""
assert len(self.tokenizers) == len(self.text_encoders)
assert len(self.text_encoders) == len(
fastvideo_args.pipeline_config.text_encoder_configs)
# Resolve selection into indices
encoder_cfgs = fastvideo_args.pipeline_config.text_encoder_configs
if encoder_index is None:
indices: list[int] = [0]
elif isinstance(encoder_index, int):
indices = [encoder_index]
else:
indices = list(encoder_index)
# validate range
num_encoders = len(self.text_encoders)
for idx in indices:
if idx < 0 or idx >= num_encoders:
raise IndexError(
f"encoder index {idx} out of range [0, {num_encoders-1}]")
# Validate indices are within range
num_encoders = len(self.text_encoders)
# Normalize input to list[str]
assert isinstance(text, str | list)
if isinstance(text, str):
texts: list[str] = [text]
else:
texts = text
embeds_list: list[torch.Tensor] = []
attn_masks_list: list[torch.Tensor] = []
preprocess_funcs = fastvideo_args.pipeline_config.preprocess_text_funcs
postprocess_funcs = fastvideo_args.pipeline_config.postprocess_text_funcs
encoder_cfgs = fastvideo_args.pipeline_config.text_encoder_configs
if return_type not in ("list", "dict", "stack"):
raise ValueError(
f"Invalid return_type '{return_type}'. Expected one of: 'list', 'dict', 'stack'"
)
target_device = device if device is not None else get_local_torch_device(
)
for i in indices:
tokenizer = self.tokenizers[i]
text_encoder = self.text_encoders[i]
encoder_config = encoder_cfgs[i]
preprocess_func = preprocess_funcs[i]
postprocess_func = postprocess_funcs[i]
processed_texts: list[str] = []
for prompt_str in texts:
processed_texts.append(preprocess_func(prompt_str))
tok_kwargs = dict(encoder_config.tokenizer_kwargs)
if max_length is not None:
tok_kwargs["max_length"] = max_length
if truncation is not None:
tok_kwargs["truncation"] = truncation
if padding is not None:
tok_kwargs["padding"] = padding
text_inputs = tokenizer(processed_texts,
**tok_kwargs).to(target_device)
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs = text_encoder(
input_ids=input_ids,
attention_mask=attention_mask,
output_hidden_states=True,
)
prompt_embeds = postprocess_func(outputs)
if dtype is not None:
prompt_embeds = prompt_embeds.to(dtype=dtype)
embeds_list.append(prompt_embeds)
if return_attention_mask:
attn_masks_list.append(attention_mask)
# Shape results according to return_type
if return_type == "list":
if return_attention_mask:
return embeds_list, attn_masks_list
return embeds_list
if return_type == "dict":
key_strs = [str(i) for i in indices]
embeds_dict = {
k: v
for k, v in zip(key_strs, embeds_list, strict=False)
}
if return_attention_mask:
attn_dict = {
k: v
for k, v in zip(key_strs, attn_masks_list, strict=False)
}
return embeds_dict, attn_dict
return embeds_dict
# return_type == "stack"
# Validate shapes are compatible
base_shape = list(embeds_list[0].shape)
for t in embeds_list[1:]:
if list(t.shape) != base_shape:
raise ValueError(
f"Cannot stack embeddings with differing shapes: {[list(t.shape) for t in embeds_list]}"
)
stacked_embeds = torch.stack(embeds_list, dim=0)
if return_attention_mask:
base_mask_shape = list(attn_masks_list[0].shape)
for m in attn_masks_list[1:]:
if list(m.shape) != base_mask_shape:
raise ValueError(
f"Cannot stack attention masks with differing shapes: {[list(m.shape) for m in attn_masks_list]}"
)
stacked_masks = torch.stack(attn_masks_list, dim=0)
return stacked_embeds, stacked_masks
return stacked_embeds
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify text encoding stage outputs."""
+14
View File
@@ -159,6 +159,20 @@ 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,6 +19,7 @@ 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
@@ -0,0 +1,108 @@
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
@@ -0,0 +1,58 @@
# 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()
+9 -1
View File
@@ -102,10 +102,18 @@ 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")
+3 -2
View File
@@ -4,7 +4,8 @@ The reference videos in the `*_reference_videos` directory are used as part of a
run `bash update_reference_videos.sh` from inside the `fastvideo/tests/ssim/` directory after running `test_inference_similarity.py` to update reference videos. Note: make sure to update the path to the corresponding device.
all reference videos are were generated on commit `4aeabbc629e0edf91477e80e795e7bb1823c71cb`
reference videos were generated on commit `4aeabbc629e0edf91477e80e795e7bb1823c71cb`
causal videos were generated on commit b318063c0a4618f1d5d99ea82ca67a06aad0d19d
## Generation Details
@@ -76,4 +77,4 @@ Wan2.1-I2V-14B-480P-Diffusers: {
### Image-to-Video Prompts
1. "An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot."
Image path: "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
Image path: "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
@@ -0,0 +1,152 @@
# SPDX-License-Identifier: Apache-2.0
import json
import os
import torch
import pytest
from fastvideo import VideoGenerator
from fastvideo.logger import init_logger
from fastvideo.tests.utils import compute_video_ssim_torchvision, write_ssim_results
from fastvideo.worker.multiproc_executor import MultiprocExecutor
logger = init_logger(__name__)
device_name = torch.cuda.get_device_name()
device_reference_folder_suffix = '_reference_videos'
if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
# Base parameters from the shell script
SF_WAN_T2V_PARAMS = {
"num_gpus": 1,
"model_path": "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
"height": 480,
"width": 832,
"num_frames": 81,
"num_inference_steps": 4,
"seed": 1024,
"sp_size": 1,
"tp_size": 1,
}
MODEL_TO_PARAMS = {
"SFWan2.1-T2V-1.3B-Diffusers": SF_WAN_T2V_PARAMS,
}
I2V_MODEL_TO_PARAMS = {
}
TEST_PROMPTS = [
"Will Smith casually eats noodles, his relaxed demeanor contrasting with the energetic background of a bustling street food market. The scene captures a mix of humor and authenticity. Mid-shot framing, vibrant lighting.",
# "A lone hiker stands atop a towering cliff, silhouetted against the vast horizon. The rugged landscape stretches endlessly beneath, its earthy tones blending into the soft blues of the sky. The scene captures the spirit of exploration and human resilience. High angle, dynamic framing, with soft natural lighting emphasizing the grandeur of nature."
]
I2V_TEST_PROMPTS = [
"An astronaut hatching from an egg, on the surface of the moon, the darkness and depth of space realised in the background. High quality, ultrarealistic detail and breath-taking movie-like camera shot.",
]
I2V_IMAGE_PATHS = [
"https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg",
]
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
@pytest.mark.parametrize("model_id", list(MODEL_TO_PARAMS.keys()))
def test_causal_similarity(prompt, ATTENTION_BACKEND, model_id):
"""
Test that runs inference with different parameters and compares the output
to reference videos using SSIM.
"""
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
script_dir = os.path.dirname(os.path.abspath(__file__))
base_output_dir = os.path.join(script_dir, 'generated_videos', model_id)
output_dir = os.path.join(base_output_dir, ATTENTION_BACKEND)
output_video_name = f"{prompt[:100]}.mp4"
os.makedirs(output_dir, exist_ok=True)
BASE_PARAMS = MODEL_TO_PARAMS[model_id]
num_inference_steps = BASE_PARAMS["num_inference_steps"]
init_kwargs = {
"num_gpus": BASE_PARAMS["num_gpus"],
"sp_size": BASE_PARAMS["sp_size"],
"tp_size": BASE_PARAMS["tp_size"],
"dit_cpu_offload": True,
}
if BASE_PARAMS.get("vae_sp"):
init_kwargs["vae_sp"] = True
init_kwargs["vae_tiling"] = True
#if "text-encoder-precision" in BASE_PARAMS:
# init_kwargs["text_encoder_precisions"] = BASE_PARAMS["text-encoder-precision"]
generation_kwargs = {
"num_inference_steps": num_inference_steps,
"output_path": output_dir,
"height": BASE_PARAMS["height"],
"width": BASE_PARAMS["width"],
"num_frames": BASE_PARAMS["num_frames"],
"seed": BASE_PARAMS["seed"],
}
if "neg_prompt" in BASE_PARAMS:
generation_kwargs["neg_prompt"] = BASE_PARAMS["neg_prompt"]
generator = VideoGenerator.from_pretrained(model_path=BASE_PARAMS["model_path"], **init_kwargs)
generator.generate_video(prompt, **generation_kwargs)
if isinstance(generator.executor, MultiprocExecutor):
generator.executor.shutdown()
assert os.path.exists(
output_dir), f"Output video was not generated at {output_dir}"
reference_folder = os.path.join(script_dir, device_reference_folder, model_id, ATTENTION_BACKEND)
if not os.path.exists(reference_folder):
logger.error("Reference folder missing")
raise FileNotFoundError(
f"Reference video folder does not exist: {reference_folder}")
# Find the matching reference video based on the prompt
reference_video_name = None
for filename in os.listdir(reference_folder):
if filename.endswith('.mp4') and prompt[:100] in filename:
reference_video_name = filename
break
if not reference_video_name:
logger.error(f"Reference video not found for prompt: {prompt} with backend: {ATTENTION_BACKEND}")
raise FileNotFoundError(f"Reference video missing")
reference_video_path = os.path.join(reference_folder, reference_video_name)
generated_video_path = os.path.join(output_dir, output_video_name)
logger.info(
f"Computing SSIM between {reference_video_path} and {generated_video_path}"
)
ssim_values = compute_video_ssim_torchvision(reference_video_path,
generated_video_path,
use_ms_ssim=True)
mean_ssim = ssim_values[0]
logger.info(f"SSIM mean value: {mean_ssim}")
logger.info(f"Writing SSIM results to directory: {output_dir}")
success = write_ssim_results(output_dir, ssim_values, reference_video_path,
generated_video_path, num_inference_steps,
prompt)
if not success:
logger.error("Failed to write SSIM results to file")
min_acceptable_ssim = 0.98
assert mean_ssim >= min_acceptable_ssim, f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} for {model_id} with backend {ATTENTION_BACKEND}"
@@ -101,7 +101,7 @@ I2V_IMAGE_PATHS = [
@pytest.mark.parametrize("prompt", I2V_TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN", "TORCH_SDPA"])
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
@pytest.mark.parametrize("model_id", list(I2V_MODEL_TO_PARAMS.keys()))
def test_i2v_inference_similarity(prompt, ATTENTION_BACKEND, model_id):
"""
@@ -0,0 +1,128 @@
import torch
import types
import pytest
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig, TextEncoderConfig, BaseEncoderOutput
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
class TensorDict(dict):
def to(self, device):
return TensorDict({k: v.to(device) for k, v in self.items()})
class FakeTokenizer:
def __call__(self, texts, **kwargs):
B = len(texts)
seq_len = int(kwargs.get("max_length", 4))
return TensorDict({
"input_ids": torch.arange(B * seq_len).view(B, seq_len),
"attention_mask": torch.ones(B, seq_len, dtype=torch.long),
})
class FakeTextEncoder(torch.nn.Module):
def __init__(self, hidden_size=8):
super().__init__()
self.hidden_size = hidden_size
def forward(self, input_ids, attention_mask, output_hidden_states=True):
B, T = input_ids.shape
last_hidden_state = torch.arange(B * T * self.hidden_size, dtype=torch.float32).view(B, T, self.hidden_size)
return types.SimpleNamespace(last_hidden_state=last_hidden_state)
def id_preprocess(x: str) -> str:
return x
def take_mean_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
# [B, T, H] -> [B, H]
return outputs.last_hidden_state.mean(dim=1)
def make_args(num_encoders=2, text_len=4, hidden_size=8):
enc_cfgs = []
preprocess_fns = []
postprocess_fns = []
for _ in range(num_encoders):
arch = TextEncoderArchConfig(text_len=text_len)
enc_cfgs.append(TextEncoderConfig(arch_config=arch))
preprocess_fns.append(id_preprocess)
postprocess_fns.append(take_mean_postprocess)
pipe_cfg = PipelineConfig(
text_encoder_configs=tuple(enc_cfgs),
text_encoder_precisions=tuple(["fp32"] * num_encoders),
preprocess_text_funcs=tuple(preprocess_fns),
postprocess_text_funcs=tuple(postprocess_fns),
)
return FastVideoArgs(model_path="", pipeline_config=pipe_cfg), hidden_size
def make_stage(num_encoders=2, hidden_size=8):
tokenizers = [FakeTokenizer() for _ in range(num_encoders)]
encoders = [FakeTextEncoder(hidden_size=hidden_size) for _ in range(num_encoders)]
return TextEncodingStage(text_encoders=encoders, tokenizers=tokenizers)
def test_encode_text_selection_and_shapes():
fastvideo_args, hidden = make_args(num_encoders=2, text_len=4, hidden_size=8)
stage = make_stage(num_encoders=2, hidden_size=hidden)
# list return, two encoders
embeds = stage.encode_text(["a", "b"], fastvideo_args, encoder_index=[0, 1])
assert isinstance(embeds, list) and len(embeds) == 2
for e in embeds:
assert e.shape == (2, hidden)
# with masks
embeds2, masks2 = stage.encode_text("a", fastvideo_args, encoder_index=[1], return_attention_mask=True)
assert len(embeds2) == 1 and len(masks2) == 1
assert embeds2[0].shape == (1, hidden)
assert masks2[0].shape == (1, 4)
# dict return
d = stage.encode_text(["a","b"], fastvideo_args, encoder_index=[0,1], return_type="dict")
assert set(d.keys()) == {"0", "1"}
assert d["0"].shape == (2, hidden)
# stack return
s = stage.encode_text(["a","b"], fastvideo_args, encoder_index=[0,1], return_type="stack")
assert s.shape == (2, 2, hidden) # [encoders, batch, hidden]
# overrides: dtype + max_length
e3, m3 = stage.encode_text(["a"], fastvideo_args, encoder_index=[0], dtype=torch.float16, return_attention_mask=True, max_length=3)
assert e3[0].dtype == torch.float16
assert m3[0].shape[1] == 3
def test_forward_integration_cfg_off_and_on():
fastvideo_args, hidden = make_args(num_encoders=2, text_len=4, hidden_size=8)
stage = make_stage(num_encoders=2, hidden_size=hidden)
# CFG off
batch = ForwardBatch(
data_type="video",
prompt="a cat",
negative_prompt="",
do_classifier_free_guidance=False,
prompt_embeds=[],
negative_prompt_embeds=None,
prompt_attention_mask=[],
negative_attention_mask=None,
)
out = stage.forward(batch, fastvideo_args)
assert len(out.prompt_embeds) == 2
for e in out.prompt_embeds:
assert e.shape[1] == hidden
# CFG on
batch2 = ForwardBatch(
data_type="video",
prompt=["a cat", "a dog"],
negative_prompt="bad picture",
do_classifier_free_guidance=True,
prompt_embeds=[],
negative_prompt_embeds=[],
prompt_attention_mask=[],
negative_attention_mask=[],
)
out2 = stage.forward(batch2, fastvideo_args)
assert len(out2.prompt_embeds) == 2
assert len(out2.negative_prompt_embeds) == 2
assert len(out2.prompt_attention_mask) == 2
assert len(out2.negative_attention_mask) == 2
@@ -0,0 +1,67 @@
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
+77 -394
View File
@@ -1,7 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
import copy
import gc
import json
import os
import time
from abc import abstractmethod
@@ -37,11 +36,10 @@ from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing)
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
EMA_FSDP, clip_grad_norm_while_handling_failing_dtensor_cases,
clip_grad_norm_while_handling_failing_dtensor_cases, count_trainable,
get_scheduler, load_distillation_checkpoint, save_distillation_checkpoint,
shift_timestep)
from fastvideo.utils import (is_vsa_available, maybe_download_model,
set_random_seed, verify_model_config_and_directory)
from fastvideo.utils import is_vsa_available, set_random_seed
import wandb # isort: skip
@@ -71,11 +69,18 @@ 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...")
@@ -84,37 +89,14 @@ 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)
if training_args.real_score_model_path:
logger.info(
f"Loading real score transformer from: {training_args.real_score_model_path}"
)
self.real_score_transformer = self.load_module_from_path(
training_args.real_score_model_path, "transformer",
training_args)
else:
self.real_score_transformer = self.get_module(
"real_score_transformer")
if training_args.fake_score_model_path:
logger.info(
f"Loading fake score transformer from: {training_args.fake_score_model_path}"
)
self.fake_score_transformer = self.load_module_from_path(
training_args.fake_score_model_path, "transformer",
training_args)
else:
self.fake_score_transformer = self.get_module(
"fake_score_transformer")
self.real_score_transformer.requires_grad_(False)
# 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.eval()
self.fake_score_transformer.requires_grad_(True)
self.fake_score_transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
@@ -137,13 +119,10 @@ class DistillationPipeline(TrainingPipeline):
if fake_score_lr == 0.0:
fake_score_lr = training_args.learning_rate
betas_str = training_args.fake_score_betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.fake_score_optimizer = torch.optim.AdamW(
fake_score_params,
lr=fake_score_lr,
betas=betas,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
@@ -171,19 +150,8 @@ class DistillationPipeline(TrainingPipeline):
self.training_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long,
device=get_local_torch_device())
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
timesteps = torch.cat((self.noise_scheduler.timesteps.cpu(),
torch.tensor([0],
dtype=torch.float32))).cuda()
self.denoising_step_list = timesteps[1000 -
self.denoising_step_list]
logger.info("Warping denoising_step_list")
self.denoising_step_list = self.denoising_step_list.to(
get_local_torch_device())
logger.info("Distillation generator model to %s denoising steps: %s",
len(self.denoising_step_list), self.denoising_step_list)
logger.info("Distillation generator model to %s denoising steps",
len(self.denoising_step_list))
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
self.min_timestep = int(self.training_args.min_timestep_ratio *
@@ -193,82 +161,6 @@ class DistillationPipeline(TrainingPipeline):
self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
self.generator_ema = None
if (self.training_args.ema_decay
is not None) and (self.training_args.ema_decay > 0.0):
self.generator_ema = EMA_FSDP(self.transformer,
decay=self.training_args.ema_decay)
logger.info(
f"Initialized generator EMA with decay={self.training_args.ema_decay}"
)
else:
logger.info("Generator EMA disabled (ema_decay <= 0.0)")
def load_module_from_path(self, model_path: str, module_type: str,
training_args: "TrainingArgs"):
"""
Load a module from a specific path using the same loading logic as the pipeline.
Args:
model_path: Path to the model
module_type: Type of module to load (e.g., "transformer")
training_args: Training arguments
Returns:
The loaded module
"""
logger.info(f"Loading {module_type} from custom path: {model_path}")
# Set flag to prevent custom weight loading for teacher/critic models
training_args._loading_teacher_critic_model = True
try:
from fastvideo.models.loader.component_loader import (
PipelineComponentLoader)
# Download the model if it's a Hugging Face model ID
local_model_path = maybe_download_model(model_path)
logger.info(f"Model downloaded/found at: {local_model_path}")
config = verify_model_config_and_directory(local_model_path)
if module_type not in config:
if hasattr(self, '_extra_config_module_map'
) and module_type in self._extra_config_module_map:
extra_module = self._extra_config_module_map[module_type]
if extra_module in config:
module_type = extra_module
logger.info(f"Using {extra_module} for {module_type}")
else:
raise ValueError(
f"Module {module_type} not found in config at {local_model_path}"
)
else:
raise ValueError(
f"Module {module_type} not found in config at {local_model_path}"
)
module_info = config[module_type]
if module_info is None:
raise ValueError(
f"Module {module_type} has null value in config at {local_model_path}"
)
transformers_or_diffusers, architecture = module_info
component_path = os.path.join(local_model_path, module_type)
module = PipelineComponentLoader.load_module(
module_name=module_type,
component_model_path=component_path,
transformers_or_diffusers=transformers_or_diffusers,
fastvideo_args=training_args,
)
logger.info(
f"Successfully loaded {module_type} from {component_path}")
return module
finally:
# Always clean up the flag
if hasattr(training_args, '_loading_teacher_critic_model'):
delattr(training_args, '_loading_teacher_critic_model')
@abstractmethod
def initialize_validation_pipeline(self, training_args: TrainingArgs):
"""Initialize validation pipeline - must be implemented by subclasses."""
@@ -278,117 +170,11 @@ 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
def apply_ema_to_model(self, model):
"""Apply EMA weights to the model for validation or inference."""
if self.generator_ema is not None:
with self.generator_ema.apply_to_model(model):
return model
return model
def get_ema_model_copy(self):
"""Get a copy of the model with EMA weights applied."""
if self.generator_ema is not None:
ema_model = copy.deepcopy(self.transformer)
self.generator_ema.copy_to_unwrapped(ema_model)
return ema_model
return None
def is_ema_ready(self, current_step: int = None):
"""Check if EMA is ready for use (after ema_start_step)."""
if current_step is None:
current_step = getattr(self, 'current_trainstep', 0)
return (self.generator_ema is not None
and current_step >= self.training_args.ema_start_step)
def save_ema_weights(self, output_dir: str, step: int):
"""Save EMA weights separately for inference purposes."""
if self.generator_ema is None:
logger.warning("Cannot save EMA weights: EMA not initialized")
return
if not self.is_ema_ready():
logger.warning(
"Cannot save EMA weights: EMA not ready yet (step < ema_start_step)"
)
return
try:
ema_model = self.get_ema_model_copy()
if ema_model is None:
logger.warning("Failed to create EMA model copy")
return
ema_save_dir = os.path.join(output_dir, f"ema_checkpoint-{step}")
os.makedirs(ema_save_dir, exist_ok=True)
# save as diffusers format
from safetensors.torch import save_file
from fastvideo.training.training_utils import (
custom_to_hf_state_dict, gather_state_dict_on_cpu_rank0)
cpu_state = gather_state_dict_on_cpu_rank0(ema_model, device=None)
if self.global_rank == 0:
weight_path = os.path.join(
ema_save_dir, "diffusion_pytorch_model.safetensors")
diffusers_state_dict = custom_to_hf_state_dict(
cpu_state, ema_model.reverse_param_names_mapping)
save_file(diffusers_state_dict, weight_path)
config_dict = ema_model.hf_config
if "dtype" in config_dict:
del config_dict["dtype"]
config_path = os.path.join(ema_save_dir, "config.json")
with open(config_path, "w") as f:
json.dump(config_dict, f, indent=4)
logger.info(f"EMA weights saved to {weight_path}")
del ema_model
except Exception as e:
logger.error(f"Failed to save EMA weights: {str(e)}")
def get_ema_stats(self):
"""Get EMA statistics for monitoring."""
if self.generator_ema is None:
return {
"ema_enabled": False,
"ema_decay": None,
"ema_start_step": self.training_args.ema_start_step,
"ema_ready": False,
"ema_step": self.current_trainstep,
}
return {
"ema_enabled": True,
"ema_decay": self.training_args.ema_decay,
"ema_start_step": self.training_args.ema_start_step,
"ema_ready": self.is_ema_ready(),
"ema_step": self.current_trainstep,
}
def reset_ema(self):
"""Reset EMA to current model weights."""
if self.generator_ema is not None:
logger.info("Resetting EMA to current model weights")
self.generator_ema.update(self.transformer)
# Force update to current weights by setting decay to 0 temporarily
original_decay = self.generator_ema.decay
self.generator_ema.decay = 0.0
self.generator_ema.update(self.transformer)
self.generator_ema.decay = original_decay
logger.info("EMA reset completed")
else:
logger.warning("Cannot reset EMA: EMA not initialized")
def _build_distill_input_kwargs(
self, noise_input: torch.Tensor, timestep: torch.Tensor,
text_dict: dict[str, torch.Tensor] | None,
@@ -802,10 +588,6 @@ class DistillationPipeline(TrainingPipeline):
self._clip_model_grad_norm_(batch_gen, self.transformer)
self.optimizer.step()
self.optimizer.zero_grad(set_to_none=True)
if self.generator_ema is not None:
self.generator_ema.update(self.transformer)
avg_dmd_loss = torch.tensor(total_dmd_loss /
gradient_accumulation_steps,
device=self.device)
@@ -856,8 +638,7 @@ class DistillationPipeline(TrainingPipeline):
self.transformer, self.fake_score_transformer, self.global_rank,
self.training_args.resume_from_checkpoint, self.optimizer,
self.fake_score_optimizer, self.train_dataloader, self.lr_scheduler,
self.fake_score_lr_scheduler, self.noise_random_generator,
self.generator_ema)
self.fake_score_lr_scheduler, self.noise_random_generator)
if resumed_step > 0:
self.init_steps = resumed_step
@@ -888,14 +669,6 @@ class DistillationPipeline(TrainingPipeline):
sum(p.numel()
for p in self.fake_score_transformer.parameters()) / 1e9)
if self.generator_ema is not None:
logger.info(" Generator EMA enabled with decay: %s",
self.training_args.ema_decay)
logger.info(" Generator EMA start step: %s",
self.training_args.ema_start_step)
else:
logger.info(" Generator EMA disabled")
@torch.no_grad()
def _log_validation(self, transformer, training_args, global_step) -> None:
training_args.inference_mode = True
@@ -927,18 +700,6 @@ class DistillationPipeline(TrainingPipeline):
transformer.eval()
# Optionally use EMA model for validation if available and ready
use_ema_for_validation = (self.training_args.use_ema
and self.is_ema_ready(global_step))
if use_ema_for_validation:
logger.info("Using EMA model for validation")
validation_transformer = self.transformer
ema_context = self.generator_ema.apply_to_model(
validation_transformer)
else:
validation_transformer = transformer
ema_context = None
validation_steps = training_args.validation_sampling_steps.split(",")
validation_steps = [int(step) for step in validation_steps]
validation_steps = [step for step in validation_steps if step > 0]
@@ -954,98 +715,50 @@ class DistillationPipeline(TrainingPipeline):
step_videos: list[np.ndarray] = []
step_captions: list[str] = []
if ema_context is not None:
with ema_context:
for validation_batch in validation_dataloader:
batch = self._prepare_validation_batch(
sampling_param, training_args, validation_batch,
num_inference_steps)
for validation_batch in validation_dataloader:
batch = self._prepare_validation_batch(sampling_param,
training_args,
validation_batch,
num_inference_steps)
negative_prompt = batch.negative_prompt
batch_negative = ForwardBatch(
data_type="video",
prompt=negative_prompt,
prompt_embeds=[],
prompt_attention_mask=[],
)
result_batch = self.validation_pipeline.prompt_encoding_stage( # type: ignore
batch_negative, training_args)
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
negative_prompt = batch.negative_prompt
batch_negative = ForwardBatch(
data_type="video",
prompt=negative_prompt,
prompt_embeds=[],
prompt_attention_mask=[],
)
result_batch = self.validation_pipeline.prompt_encoding_stage( # type: ignore
batch_negative, training_args)
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
logger.info(
"rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
logger.info("rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
self.global_rank,
self.rank_in_sp_group,
batch.prompt,
local_main_process_only=False)
assert batch.prompt is not None and isinstance(
batch.prompt, str)
step_captions.append(batch.prompt)
assert batch.prompt is not None and isinstance(
batch.prompt, str)
step_captions.append(batch.prompt)
# Run validation inference
with torch.no_grad():
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
if self.rank_in_sp_group != 0:
continue
# Run validation inference
with torch.no_grad():
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
if self.rank_in_sp_group != 0:
continue
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
step_videos.append(frames)
else:
# Use original transformer without EMA
for validation_batch in validation_dataloader:
batch = self._prepare_validation_batch(
sampling_param, training_args, validation_batch,
num_inference_steps)
negative_prompt = batch.negative_prompt
batch_negative = ForwardBatch(
data_type="video",
prompt=negative_prompt,
prompt_embeds=[],
prompt_attention_mask=[],
)
result_batch = self.validation_pipeline.prompt_encoding_stage( # type: ignore
batch_negative, training_args)
self.negative_prompt_embeds, self.negative_prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
logger.info(
"rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
self.global_rank,
self.rank_in_sp_group,
batch.prompt,
local_main_process_only=False)
assert batch.prompt is not None and isinstance(
batch.prompt, str)
step_captions.append(batch.prompt)
# Run validation inference
with torch.no_grad():
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
if self.rank_in_sp_group != 0:
continue
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
step_videos.append(frames)
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
step_videos.append(frames)
# Log validation results for this step
world_group = get_world_group()
@@ -1122,16 +835,16 @@ class DistillationPipeline(TrainingPipeline):
latents.dtype)
else:
latents += self.vae.shift_factor
with torch.autocast("cuda", dtype=torch.bfloat16):
video = self.vae.decode(latents)
video = (video / 2 + 0.5).clamp(0, 1)
video = video.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[latent_key] = wandb.Video(
video, fps=24, format="mp4") # change to 16 for Wan2.1
# Clean up references
del video, latents
with torch.autocast("cuda", dtype=torch.bfloat16):
video = self.vae.decode(latents)
video = (video / 2 + 0.5).clamp(0, 1)
video = video.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[latent_key] = wandb.Video(
video, fps=24, format="mp4") # change to 16 for Wan2.1
# Clean up references
del video, latents
# Process DMD training data if available - use decode_stage instead of self.vae.decode
if 'generator_pred_video' in dmd_latents_vis_dict:
@@ -1183,6 +896,14 @@ 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)
@@ -1192,10 +913,6 @@ class DistillationPipeline(TrainingPipeline):
device="cpu").manual_seed(self.seed)
logger.info("Initialized random seeds with seed: %s", seed)
# Initialize current_trainstep for EMA ready checks
#TODO: check if needed
self.current_trainstep = self.init_steps
# Resume from checkpoint if specified (this will restore random states)
if self.training_args.resume_from_checkpoint:
self._resume_from_checkpoint()
@@ -1239,14 +956,6 @@ class DistillationPipeline(TrainingPipeline):
self.current_trainstep = step
training_batch.current_vsa_sparsity = current_vsa_sparsity
if (step >= self.training_args.ema_start_step) and \
(self.generator_ema is None) and (self.training_args.ema_decay > 0):
self.generator_ema = EMA_FSDP(
self.transformer, decay=self.training_args.ema_decay)
logger.info(
f"Created generator EMA at step {step} with decay={self.training_args.ema_decay}"
)
with torch.autocast("cuda", dtype=torch.bfloat16):
training_batch = self.train_one_step(training_batch)
@@ -1260,19 +969,11 @@ class DistillationPipeline(TrainingPipeline):
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"total_loss":
f"{total_loss:.4f}",
"generator_loss":
f"{generator_loss:.4f}",
"fake_score_loss":
f"{fake_score_loss:.4f}",
"step_time":
f"{step_time:.2f}s",
"grad_norm":
grad_norm,
"ema":
"✓" if (self.generator_ema is not None and self.is_ema_ready())
else "✗",
"total_loss": f"{total_loss:.4f}",
"generator_loss": f"{generator_loss:.4f}",
"fake_score_loss": f"{fake_score_loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
})
progress_bar.update(1)
@@ -1300,15 +1001,6 @@ class DistillationPipeline(TrainingPipeline):
if use_vsa:
log_data["VSA_train_sparsity"] = current_vsa_sparsity
if self.generator_ema is not None:
log_data["ema_enabled"] = True
log_data["ema_decay"] = self.training_args.ema_decay
else:
log_data["ema_enabled"] = False
ema_stats = self.get_ema_stats()
log_data.update(ema_stats)
if training_batch.dmd_latent_vis_dict:
dmd_additional_logs = {
"generator_timestep":
@@ -1340,8 +1032,7 @@ class DistillationPipeline(TrainingPipeline):
self.global_rank, self.training_args.output_dir, step,
self.optimizer, self.fake_score_optimizer,
self.train_dataloader, self.lr_scheduler,
self.fake_score_lr_scheduler, self.noise_random_generator,
self.generator_ema)
self.fake_score_lr_scheduler, self.noise_random_generator)
if self.transformer:
self.transformer.train()
@@ -1358,11 +1049,7 @@ class DistillationPipeline(TrainingPipeline):
self.global_rank,
self.training_args.output_dir,
f"{step}_weight_only",
only_save_generator_weight=True,
generator_ema=self.generator_ema)
if self.training_args.use_ema and self.is_ema_ready():
self.save_ema_weights(self.training_args.output_dir, step)
only_save_generator_weight=True)
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
if self.training_args.log_visualization:
@@ -1382,11 +1069,7 @@ class DistillationPipeline(TrainingPipeline):
self.training_args.output_dir, self.training_args.max_train_steps,
self.optimizer, self.fake_score_optimizer, self.train_dataloader,
self.lr_scheduler, self.fake_score_lr_scheduler,
self.noise_random_generator, self.generator_ema)
if self.training_args.use_ema and self.is_ema_ready():
self.save_ema_weights(self.training_args.output_dir,
self.training_args.max_train_steps)
self.noise_random_generator)
if get_sp_group():
cleanup_dist_env_and_memory()
@@ -1,901 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import copy
import gc
import logging
from typing import Any
from collections import deque
import time
import torch
import torch.nn.functional as F
import wandb
from tqdm.auto import tqdm
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.pipelines import TrainingBatch
from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases,
EMA_FSDP,
save_distillation_checkpoint,
)
from fastvideo.models.utils import pred_noise_to_pred_video
from fastvideo.distributed import get_world_group
import torch.distributed as dist
import numpy as np
from fastvideo.utils import set_random_seed, is_vsa_available
import fastvideo.envs as envs
from einops import rearrange
logger = init_logger(__name__)
vsa_available = is_vsa_available()
class SelfForcingDistillationPipeline(DistillationPipeline):
"""
A self-forcing distillation pipeline that alternates between training
the generator and critic based on the self-forcing methodology.
This implementation follows the self-forcing approach where:
1. Generator and critic are trained in alternating steps
2. Generator loss uses DMD-style loss with the critic as fake score
3. Critic loss trains the fake score model to distinguish real vs fake
"""
def initialize_training_pipeline(self, training_args: TrainingArgs):
"""Initialize the self-forcing training pipeline."""
logger.info("Initializing self-forcing distillation pipeline...")
super().initialize_training_pipeline(training_args)
self.dfake_gen_update_ratio = getattr(training_args, 'dfake_gen_update_ratio', 5)
logger.info(f"Self-forcing generator update ratio: {self.dfake_gen_update_ratio}")
def generator_loss(self, training_batch: TrainingBatch) -> tuple[torch.Tensor, dict[str, Any]]:
"""
Compute generator loss using DMD-style approach.
The generator tries to fool the critic (fake_score_transformer).
"""
with set_forward_context(
current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata_vsa):
if self.training_args.simulate_generator_forward:
generator_pred_video = self._generator_multi_step_simulation_forward(training_batch)
else:
generator_pred_video = self._generator_forward(training_batch)
with set_forward_context(
current_timestep=training_batch.timesteps,
attn_metadata=training_batch.attn_metadata):
dmd_loss = self._dmd_forward(
generator_pred_video=generator_pred_video,
training_batch=training_batch
)
log_dict = {
"dmdtrain_gradient_norm": torch.tensor(0.0, device=self.device)
}
return dmd_loss, log_dict
def critic_loss(self, training_batch: TrainingBatch) -> tuple[torch.Tensor, dict[str, Any]]:
"""
Compute critic loss using flow matching between noise and generator output.
The critic learns to predict the flow from noise to the generator's output.
"""
updated_batch, flow_matching_loss = self.faker_score_forward(training_batch)
training_batch.fake_score_latent_vis_dict = updated_batch.fake_score_latent_vis_dict
log_dict = {}
return flow_matching_loss, log_dict
def _generator_forward(self, training_batch: TrainingBatch) -> torch.Tensor:
"""Forward pass through generator with KV cache support for causal generation."""
latents = training_batch.latents
dtype = latents.dtype
batch_size = latents.shape[0]
# Step 1: Sample a timestep from denoising_step_list
index = torch.randint(0,
len(self.denoising_step_list), [1],
device=self.device,
dtype=torch.long)
timestep = self.denoising_step_list[index]
training_batch.dmd_latent_vis_dict["generator_timestep"] = timestep
# Step 2: Initialize KV cache and cross-attention cache for causal generation
kv_cache, crossattn_cache = self._initialize_simulation_caches(batch_size, dtype, self.device)
if getattr(self.training_args, 'validate_cache_structure', False):
self._validate_cache_structure(kv_cache, crossattn_cache, batch_size)
# Step 3: Add noise to latents
noise = torch.randn(self.video_latent_shape,
device=self.device,
dtype=dtype)
if self.sp_world_size > 1:
noise = rearrange(noise,
"b (n t) c h w -> b n t c h w",
n=self.sp_world_size).contiguous()
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
noisy_latent = self.noise_scheduler.add_noise(latents.flatten(0, 1),
noise.flatten(0, 1),
timestep).unflatten(
0,
(1, latents.shape[1]))
# Step 4: Build input kwargs with KV cache support
training_batch = self._build_distill_input_kwargs(
noisy_latent, timestep, training_batch.conditional_dict,
training_batch)
# Step 5: Forward pass with KV cache if available
if hasattr(self.transformer, '_forward_inference'):
# Use causal inference forward with KV cache
pred_noise = self.transformer(
hidden_states=training_batch.input_kwargs['hidden_states'],
encoder_hidden_states=training_batch.input_kwargs['encoder_hidden_states'],
timestep=training_batch.input_kwargs['timestep'],
encoder_hidden_states_image=training_batch.input_kwargs.get('encoder_hidden_states_image'),
kv_cache=kv_cache,
crossattn_cache=crossattn_cache,
current_start=0, # Start from beginning for single-step
cache_start=0
).permute(0, 2, 1, 3, 4)
else:
# Fallback to regular forward
pred_noise = self.transformer(**training_batch.input_kwargs).permute(
0, 2, 1, 3, 4)
# Step 6: Convert noise prediction to video prediction
pred_video = pred_noise_to_pred_video(
pred_noise=pred_noise.flatten(0, 1),
noise_input_latent=noisy_latent.flatten(0, 1),
timestep=timestep,
scheduler=self.noise_scheduler).unflatten(0, pred_noise.shape[:2])
self._reset_simulation_caches(kv_cache, crossattn_cache)
return pred_video
def _generator_multi_step_simulation_forward(
self, training_batch: TrainingBatch) -> torch.Tensor:
"""Forward pass through student transformer matching inference procedure with KV cache management."""
latents = training_batch.latents
dtype = latents.dtype
batch_size = latents.shape[0]
# Step 1: Randomly sample a target timestep index from denoising_step_list
target_timestep_idx = torch.randint(0,
len(self.denoising_step_list), [1],
device=self.device,
dtype=torch.long)
target_timestep_idx_int = target_timestep_idx.item()
target_timestep = self.denoising_step_list[target_timestep_idx]
# Step 2: Initialize KV cache and cross-attention cache for causal generation
kv_cache, crossattn_cache = self._initialize_simulation_caches(batch_size, dtype, self.device)
# Validate cache structure (can be disabled in production)
if getattr(self.training_args, 'validate_cache_structure', False):
self._validate_cache_structure(kv_cache, crossattn_cache, batch_size)
# Step 3: Simulate the multi-step inference process up to the target timestep
# Start from pure noise like in inference
current_noise_latents = torch.randn(self.video_latent_shape,
device=self.device,
dtype=dtype)
if self.sp_world_size > 1:
current_noise_latents = rearrange(
current_noise_latents,
"b (n t) c h w -> b n t c h w",
n=self.sp_world_size).contiguous()
current_noise_latents = current_noise_latents[:, self.
rank_in_sp_group, :, :, :, :]
current_noise_latents_copy = current_noise_latents.clone()
# Step 4: Determine gradient masking for frame-level selective training
num_frames = current_noise_latents.shape[1]
num_frame_per_block = getattr(self.training_args, 'num_frame_per_block', 3)
independent_first_frame = getattr(self.training_args, 'independent_first_frame', False)
enable_gradient_masking = getattr(self.training_args, 'enable_gradient_masking', True)
gradient_mask_last_n_frames = getattr(self.training_args, 'gradient_mask_last_n_frames', 21)
gradient_mask = None
if enable_gradient_masking:
# Calculate which frames should have gradients enabled (last N frames)
start_gradient_frame_index = max(0, num_frames - gradient_mask_last_n_frames)
gradient_mask = torch.ones(batch_size, num_frames, dtype=torch.bool, device=self.device)
if independent_first_frame:
# First frame is independent, disable gradients for early frames
gradient_mask[:, :max(1, start_gradient_frame_index)] = False
else:
# Disable gradients for early blocks
num_early_blocks = start_gradient_frame_index // num_frame_per_block
gradient_mask[:, :num_early_blocks * num_frame_per_block] = False
# Step 5: Multi-step simulation with causal block processing
num_frames_per_block = getattr(self.training_args, 'num_frame_per_block', 3)
t = num_frames
if not independent_first_frame or (independent_first_frame and hasattr(training_batch, 'image_latent') and training_batch.image_latent is not None):
if t % num_frames_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frames_per_block for causal DMD denoising"
)
num_blocks = t // num_frames_per_block
block_sizes = [num_frames_per_block] * num_blocks
else:
if (t - 1) % num_frames_per_block != 0:
raise ValueError(
"(num_frames - 1) must be divisible by num_frame_per_block when independent_first_frame=True"
)
num_blocks = (t - 1) // num_frames_per_block
block_sizes = [1] + [num_frames_per_block] * num_blocks
max_target_idx = len(self.denoising_step_list) - 1
noise_latents = []
noise_latent_index = target_timestep_idx_int - 1
if max_target_idx > 0:
# Run student model for all steps before the target timestep with causal blocks
with torch.no_grad():
start_index = 0
current_start_frame = 0
# Process each causal block
for current_num_frames in block_sizes:
# Extract current block from the full latents
block_latents = current_noise_latents[:, start_index:start_index + current_num_frames, :, :, :]
# Process denoising timesteps for this block
for step_idx in range(max_target_idx):
current_timestep = self.denoising_step_list[step_idx]
current_timestep_tensor = current_timestep * torch.ones(
1, device=self.device, dtype=torch.long)
# Build input kwargs with KV cache support for this block
training_batch_temp = self._build_distill_input_kwargs(
block_latents, current_timestep_tensor,
training_batch.conditional_dict, training_batch)
# Add KV cache parameters for causal generation
if hasattr(self.transformer, '_forward_inference'):
# Use causal inference forward with KV cache
pred_flow = self.transformer(
hidden_states=training_batch_temp.input_kwargs['hidden_states'],
encoder_hidden_states=training_batch_temp.input_kwargs['encoder_hidden_states'],
timestep=training_batch_temp.input_kwargs['timestep'],
encoder_hidden_states_image=training_batch_temp.input_kwargs.get('encoder_hidden_states_image'),
kv_cache=kv_cache,
crossattn_cache=crossattn_cache,
current_start=current_start_frame * 1560, # TODO: remove hardcode
cache_start=0
).permute(0, 2, 1, 3, 4)
else:
# Fallback to regular forward
pred_flow = self.transformer(**training_batch_temp.input_kwargs).permute(0, 2, 1, 3, 4)
pred_clean = pred_noise_to_pred_video(
pred_noise=pred_flow.flatten(0, 1),
noise_input_latent=block_latents.flatten(0, 1),
timestep=current_timestep_tensor,
scheduler=self.noise_scheduler).unflatten(
0, pred_flow.shape[:2])
# Add noise for the next timestep
if step_idx < max_target_idx - 1:
next_timestep = self.denoising_step_list[step_idx + 1]
next_timestep_tensor = next_timestep * torch.ones(
1, device=self.device, dtype=torch.long)
# Generate noise for this block size
block_noise_shape = (batch_size, current_num_frames, *self.video_latent_shape[2:])
noise = torch.randn(block_noise_shape,
device=self.device,
dtype=pred_clean.dtype)
if self.sp_world_size > 1:
noise = rearrange(noise,
"b (n t) c h w -> b n t c h w",
n=self.sp_world_size).contiguous()
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
block_latents = self.noise_scheduler.add_noise(
pred_clean.flatten(0, 1), noise.flatten(0, 1),
next_timestep_tensor).unflatten(0, pred_clean.shape[:2])
else:
# Final step: use clean prediction
block_latents = pred_clean
# Store the processed block result
latent_copy = block_latents.clone()
noise_latents.append(latent_copy)
# Update indices for next block
current_start_frame += current_num_frames
start_index += current_num_frames
# Reconstruct full latents from blocks
if noise_latents:
current_noise_latents = torch.cat(noise_latents, dim=1)
# Step 6: Use the simulated noisy input for the final training step
if noise_latent_index >= 0:
assert noise_latent_index < len(
self.denoising_step_list
) - 1, "noise_latent_index is out of bounds"
noisy_input = current_noise_latents
else:
noisy_input = current_noise_latents_copy
# Step 7: Final student prediction with selective gradient computation
training_batch = self._build_distill_input_kwargs(
noisy_input, target_timestep, training_batch.conditional_dict,
training_batch)
# Apply gradient masking during final forward pass
if gradient_mask is not None:
# Create a custom forward function that applies gradient masking
def masked_forward():
pred_flow = self.transformer(**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
pred_video = pred_noise_to_pred_video(
pred_noise=pred_flow.flatten(0, 1),
noise_input_latent=noisy_input.flatten(0, 1),
timestep=target_timestep,
scheduler=self.noise_scheduler).unflatten(0, pred_flow.shape[:2])
# Apply gradient masking: detach frames that shouldn't contribute gradients
masked_pred_video = pred_video.clone()
masked_pred_video = torch.where(
gradient_mask.unsqueeze(-1).unsqueeze(-1).unsqueeze(-1), # Broadcast to [B, F, 1, 1, 1]
pred_video, # Keep original values where gradient_mask is True
pred_video.detach() # Detach where gradient_mask is False
)
return masked_pred_video
pred_video = masked_forward()
else:
pred_flow = self.transformer(**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
pred_video = pred_noise_to_pred_video(
pred_noise=pred_flow.flatten(0, 1),
noise_input_latent=noisy_input.flatten(0, 1),
timestep=target_timestep,
scheduler=self.noise_scheduler).unflatten(0, pred_flow.shape[:2])
training_batch.dmd_latent_vis_dict[
"generator_timestep"] = target_timestep.float().detach()
# Store gradient mask information for debugging
if gradient_mask is not None:
training_batch.dmd_latent_vis_dict["gradient_mask"] = gradient_mask.float()
start_gradient_frame_index = max(0, num_frames - gradient_mask_last_n_frames)
training_batch.dmd_latent_vis_dict["start_gradient_frame_index"] = torch.tensor(
start_gradient_frame_index, dtype=torch.float32, device=self.device)
self._reset_simulation_caches(kv_cache, crossattn_cache)
return pred_video
def _initialize_simulation_caches(self, batch_size: int, dtype: torch.dtype, device: torch.device):
"""Initialize KV cache and cross-attention cache for multi-step simulation."""
num_transformer_blocks = len(self.transformer.blocks)
# Calculate frame sequence length based on input dimensions and patch size
# From the training batch, we can get the actual latent dimensions
latent_shape = self.video_latent_shape_sp # This is set in _prepare_dit_inputs
batch_size_actual, num_frames, num_channels, height, width = latent_shape
# Get patch size from transformer config
p_t, p_h, p_w = self.transformer.patch_size
post_patch_height = height // p_h
post_patch_width = width // p_w
# Frame sequence length is the spatial sequence length per frame
frame_seq_length = post_patch_height * post_patch_width
# Get local attention size from transformer config
local_attn_size = getattr(self.transformer, 'local_attn_size', -1)
# Get model configuration parameters - handle FSDP wrapping
if hasattr(self.transformer, 'config'):
config = self.transformer.config
num_attention_heads = config.num_attention_heads
attention_head_dim = config.attention_head_dim
text_len = config.text_len
else:
# Fallback to direct attribute access for non-FSDP models
num_attention_heads = getattr(self.transformer, 'num_attention_heads', 40)
attention_head_dim = getattr(self.transformer, 'attention_head_dim', 128)
text_len = getattr(self.transformer, 'text_len', 512)
num_max_frames = getattr(self.training_args, "num_frames", num_frames)
kv_cache_size = num_max_frames * frame_seq_length
kv_cache = []
for _ in range(num_transformer_blocks):
kv_cache.append({
"k": torch.zeros([batch_size, kv_cache_size, num_attention_heads, attention_head_dim], dtype=dtype, device=device),
"v": torch.zeros([batch_size, kv_cache_size, num_attention_heads, attention_head_dim], dtype=dtype, device=device),
"global_end_index": torch.tensor([0], dtype=torch.long, device=device),
"local_end_index": torch.tensor([0], dtype=torch.long, device=device)
})
# Initialize cross-attention cache
crossattn_cache = []
for _ in range(num_transformer_blocks):
crossattn_cache.append({
"k": torch.zeros([batch_size, text_len, num_attention_heads, attention_head_dim], dtype=dtype, device=device),
"v": torch.zeros([batch_size, text_len, num_attention_heads, attention_head_dim], dtype=dtype, device=device),
"is_init": False
})
return kv_cache, crossattn_cache
def _reset_simulation_caches(self, kv_cache, crossattn_cache):
"""Reset KV cache and cross-attention cache to clean state."""
if kv_cache is not None:
for cache_dict in kv_cache:
cache_dict["global_end_index"].fill_(0)
cache_dict["local_end_index"].fill_(0)
cache_dict["k"].zero_()
cache_dict["v"].zero_()
if crossattn_cache is not None:
for cache_dict in crossattn_cache:
cache_dict["is_init"] = False
cache_dict["k"].zero_()
cache_dict["v"].zero_()
def _validate_cache_structure(self, kv_cache, crossattn_cache, batch_size: int):
"""Validate that cache structures are correctly initialized."""
num_transformer_blocks = len(self.transformer.blocks)
# Get model configuration parameters - handle FSDP wrapping
if hasattr(self.transformer, 'config'):
config = self.transformer.config
num_attention_heads = config.num_attention_heads
attention_head_dim = config.attention_head_dim
text_len = config.text_len
else:
# Fallback to direct attribute access for non-FSDP models
num_attention_heads = getattr(self.transformer, 'num_attention_heads', 40)
attention_head_dim = getattr(self.transformer, 'attention_head_dim', 128)
text_len = getattr(self.transformer, 'text_len', 512)
if kv_cache is not None:
assert len(kv_cache) == num_transformer_blocks, f"Expected {num_transformer_blocks} transformer blocks, got {len(kv_cache)}"
for i, cache_dict in enumerate(kv_cache):
assert "k" in cache_dict and "v" in cache_dict, f"Missing k/v in kv_cache block {i}"
assert "global_end_index" in cache_dict and "local_end_index" in cache_dict, f"Missing indices in kv_cache block {i}"
assert cache_dict["k"].shape[0] == batch_size, f"Batch size mismatch in kv_cache block {i}"
assert cache_dict["v"].shape[0] == batch_size, f"Batch size mismatch in kv_cache block {i}"
assert cache_dict["k"].shape[2] == num_attention_heads, f"Attention heads mismatch in kv_cache block {i}"
assert cache_dict["k"].shape[3] == attention_head_dim, f"Attention head dim mismatch in kv_cache block {i}"
if crossattn_cache is not None:
assert len(crossattn_cache) == num_transformer_blocks, f"Expected {num_transformer_blocks} transformer blocks, got {len(crossattn_cache)}"
for i, cache_dict in enumerate(crossattn_cache):
assert "k" in cache_dict and "v" in cache_dict, f"Missing k/v in crossattn_cache block {i}"
assert "is_init" in cache_dict, f"Missing is_init in crossattn_cache block {i}"
assert cache_dict["k"].shape[0] == batch_size, f"Batch size mismatch in crossattn_cache block {i}"
assert cache_dict["v"].shape[0] == batch_size, f"Batch size mismatch in crossattn_cache block {i}"
assert cache_dict["k"].shape[1] == text_len, f"Text length mismatch in crossattn_cache block {i}"
assert cache_dict["k"].shape[2] == num_attention_heads, f"Attention heads mismatch in crossattn_cache block {i}"
assert cache_dict["k"].shape[3] == attention_head_dim, f"Attention head dim mismatch in crossattn_cache block {i}"
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
"""
Self-forcing training step that alternates between generator and critic training.
"""
gradient_accumulation_steps = getattr(self.training_args, 'gradient_accumulation_steps', 1)
train_generator = (self.current_trainstep % self.dfake_gen_update_ratio == 0)
batches = []
for _ in range(gradient_accumulation_steps):
batch = self._prepare_distillation(training_batch)
batch = self._get_next_batch(batch)
batch = self._normalize_dit_input(batch)
batch = self._prepare_dit_inputs(batch)
batch = self._build_attention_metadata(batch)
batch.attn_metadata_vsa = copy.deepcopy(batch.attn_metadata)
if batch.attn_metadata is not None:
batch.attn_metadata.VSA_sparsity = 0.0
batches.append(batch)
training_batch.dmd_latent_vis_dict = {}
training_batch.fake_score_latent_vis_dict = {}
if train_generator:
logger.debug(f"Training generator at step {self.current_trainstep}")
self.optimizer.zero_grad()
total_generator_loss = 0.0
generator_log_dict = {}
for batch in batches:
# Create a new batch with detached tensors
batch_gen = TrainingBatch()
for key, value in batch.__dict__.items():
if isinstance(value, torch.Tensor):
setattr(batch_gen, key, value.detach().clone())
elif isinstance(value, dict):
setattr(batch_gen, key, {k: v.detach().clone() if isinstance(v, torch.Tensor) else copy.deepcopy(v) for k, v in value.items()})
else:
setattr(batch_gen, key, copy.deepcopy(value))
generator_loss, gen_log_dict = self.generator_loss(batch_gen)
with set_forward_context(
current_timestep=batch_gen.timesteps,
attn_metadata=batch_gen.attn_metadata):
(generator_loss / gradient_accumulation_steps).backward()
total_generator_loss += generator_loss.detach().item()
generator_log_dict.update(gen_log_dict)
# Store visualization data from generator training
if hasattr(batch_gen, 'dmd_latent_vis_dict'):
training_batch.dmd_latent_vis_dict.update(batch_gen.dmd_latent_vis_dict)
self._clip_model_grad_norm_(batch_gen, self.transformer)
self.optimizer.step()
self.lr_scheduler.step()
if self.generator_ema is not None:
self.generator_ema.update(self.transformer)
avg_generator_loss = torch.tensor(
total_generator_loss / gradient_accumulation_steps,
device=self.device
)
world_group = get_world_group()
world_group.all_reduce(avg_generator_loss, op=torch.distributed.ReduceOp.AVG)
training_batch.generator_loss = avg_generator_loss.item()
else:
training_batch.generator_loss = 0.0
logger.debug(f"Training critic at step {self.current_trainstep}")
self.fake_score_optimizer.zero_grad()
total_critic_loss = 0.0
critic_log_dict = {}
for batch in batches:
# Create a new batch with detached tensors
batch_critic = TrainingBatch()
for key, value in batch.__dict__.items():
if isinstance(value, torch.Tensor):
setattr(batch_critic, key, value.detach().clone())
elif isinstance(value, dict):
setattr(batch_critic, key, {k: v.detach().clone() if isinstance(v, torch.Tensor) else copy.deepcopy(v) for k, v in value.items()})
else:
setattr(batch_critic, key, copy.deepcopy(value))
critic_loss, crit_log_dict = self.critic_loss(batch_critic)
with set_forward_context(
current_timestep=batch_critic.timesteps,
attn_metadata=batch_critic.attn_metadata):
(critic_loss / gradient_accumulation_steps).backward()
total_critic_loss += critic_loss.detach().item()
critic_log_dict.update(crit_log_dict)
# Store visualization data from critic training
if hasattr(batch_critic, 'fake_score_latent_vis_dict'):
training_batch.fake_score_latent_vis_dict.update(batch_critic.fake_score_latent_vis_dict)
self._clip_model_grad_norm_(batch_critic, self.fake_score_transformer)
self.fake_score_optimizer.step()
self.fake_score_lr_scheduler.step()
avg_critic_loss = torch.tensor(
total_critic_loss / gradient_accumulation_steps,
device=self.device
)
world_group = get_world_group()
world_group.all_reduce(avg_critic_loss, op=torch.distributed.ReduceOp.AVG)
training_batch.fake_score_loss = avg_critic_loss.item()
training_batch.total_loss = training_batch.generator_loss + training_batch.fake_score_loss
return training_batch
def _log_training_info(self) -> None:
"""Log self-forcing specific training information."""
super()._log_training_info()
logger.info("Self-forcing specific settings:")
logger.info(" Generator update ratio: %s", self.dfake_gen_update_ratio)
def visualize_intermediate_latents(self, training_batch: TrainingBatch,
training_args: TrainingArgs, step: int):
"""Add visualization data to wandb logging and save frames to disk."""
wandb_loss_dict = {}
# Debug logging
logger.info(f"Step {step}: Starting visualization")
if hasattr(training_batch, 'dmd_latent_vis_dict'):
logger.info(f"DMD latent keys: {list(training_batch.dmd_latent_vis_dict.keys())}")
if hasattr(training_batch, 'fake_score_latent_vis_dict'):
logger.info(f"Fake score latent keys: {list(training_batch.fake_score_latent_vis_dict.keys())}")
# Process generator predictions if available
if hasattr(training_batch, 'dmd_latent_vis_dict') and training_batch.dmd_latent_vis_dict:
dmd_latents_vis_dict = training_batch.dmd_latent_vis_dict
dmd_log_keys = ['generator_pred_video', 'real_score_pred_video', 'faker_score_pred_video']
for latent_key in dmd_log_keys:
if latent_key in dmd_latents_vis_dict:
logger.info(f"Processing DMD latent: {latent_key}")
latents = dmd_latents_vis_dict[latent_key]
if not isinstance(latents, torch.Tensor):
logger.warning(f"Expected tensor for {latent_key}, got {type(latents)}")
continue
latents = latents.detach()
latents = latents.permute(0, 2, 1, 3, 4)
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
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
try:
with torch.autocast("cuda", dtype=torch.bfloat16):
video = self.vae.decode(latents)
video = (video / 2 + 0.5).clamp(0, 1)
video = video.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[f"dmd_{latent_key}"] = wandb.Video(video, fps=24, format="mp4")
logger.info(f"Successfully processed DMD latent: {latent_key}")
except Exception as e:
logger.error(f"Error processing DMD latent {latent_key}: {str(e)}")
del video, latents
# Process critic predictions
if hasattr(training_batch, 'fake_score_latent_vis_dict') and training_batch.fake_score_latent_vis_dict:
fake_score_latents_vis_dict = training_batch.fake_score_latent_vis_dict
fake_score_log_keys = ['generator_pred_video']
for latent_key in fake_score_log_keys:
if latent_key in fake_score_latents_vis_dict:
logger.info(f"Processing critic latent: {latent_key}")
latents = fake_score_latents_vis_dict[latent_key]
if not isinstance(latents, torch.Tensor):
logger.warning(f"Expected tensor for {latent_key}, got {type(latents)}")
continue
latents = latents.detach()
latents = latents.permute(0, 2, 1, 3, 4)
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
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
try:
with torch.autocast("cuda", dtype=torch.bfloat16):
video = self.vae.decode(latents)
video = (video / 2 + 0.5).clamp(0, 1)
video = video.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
video = (video * 255).numpy().astype(np.uint8)
wandb_loss_dict[f"critic_{latent_key}"] = wandb.Video(video, fps=24, format="mp4")
logger.info(f"Successfully processed critic latent: {latent_key}")
except Exception as e:
logger.error(f"Error processing critic latent {latent_key}: {str(e)}")
del video, latents
# Log metadata
if hasattr(training_batch, 'dmd_latent_vis_dict') and training_batch.dmd_latent_vis_dict:
if "generator_timestep" in training_batch.dmd_latent_vis_dict:
wandb_loss_dict["generator_timestep"] = training_batch.dmd_latent_vis_dict["generator_timestep"].item()
if "dmd_timestep" in training_batch.dmd_latent_vis_dict:
wandb_loss_dict["dmd_timestep"] = training_batch.dmd_latent_vis_dict["dmd_timestep"].item()
if hasattr(training_batch, 'fake_score_latent_vis_dict') and training_batch.fake_score_latent_vis_dict:
if "fake_score_timestep" in training_batch.fake_score_latent_vis_dict:
wandb_loss_dict["fake_score_timestep"] = training_batch.fake_score_latent_vis_dict["fake_score_timestep"].item()
# Log final dict contents
logger.info(f"Final wandb_loss_dict keys: {list(wandb_loss_dict.keys())}")
if self.global_rank == 0:
wandb.log(wandb_loss_dict, step=step)
def train(self) -> None:
"""Main training loop with self-forcing specific logging."""
assert self.training_args.seed is not None, "seed must be set"
seed = self.training_args.seed
# Set the same seed within each SP group to ensure reproducibility
if self.sp_world_size > 1:
# Use the same seed for all processes within the same SP group
sp_group_seed = seed + (self.global_rank // self.sp_world_size)
set_random_seed(sp_group_seed)
logger.info("Rank %s: Using SP group seed %s", self.global_rank,
sp_group_seed)
else:
set_random_seed(seed + self.global_rank)
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
self.noise_gen_cuda = torch.Generator(device="cuda").manual_seed(
self.seed)
self.validation_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed)
logger.info("Initialized random seeds with seed: %s", seed)
self.current_trainstep = self.init_steps
if self.training_args.resume_from_checkpoint:
self._resume_from_checkpoint()
logger.info("Resumed from checkpoint, random states restored")
else:
logger.info("Starting training from scratch")
self.train_loader_iter = iter(self.train_dataloader)
step_times = deque(maxlen=100)
self._log_training_info()
self._log_validation(self.transformer, self.training_args,
self.init_steps)
progress_bar = tqdm(
range(0, self.training_args.max_train_steps),
initial=self.init_steps,
desc="Steps",
disable=self.local_rank > 0,
)
use_vsa = vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN"
for step in range(self.init_steps + 1,
self.training_args.max_train_steps + 1):
start_time = time.perf_counter()
if use_vsa:
vsa_sparsity = self.training_args.VSA_sparsity
vsa_decay_rate = self.training_args.VSA_decay_rate
vsa_decay_interval_steps = self.training_args.VSA_decay_interval_steps
if vsa_decay_interval_steps > 1:
current_decay_times = min(step // vsa_decay_interval_steps,
vsa_sparsity // vsa_decay_rate)
current_vsa_sparsity = current_decay_times * vsa_decay_rate
else:
current_vsa_sparsity = vsa_sparsity
else:
current_vsa_sparsity = 0.0
training_batch = TrainingBatch()
self.current_trainstep = step
training_batch.current_vsa_sparsity = current_vsa_sparsity
if (step >= self.training_args.ema_start_step) and \
(self.generator_ema is None) and (self.training_args.ema_decay > 0):
self.generator_ema = EMA_FSDP(self.transformer, decay=self.training_args.ema_decay)
logger.info(f"Created generator EMA at step {step} with decay={self.training_args.ema_decay}")
with torch.autocast("cuda", dtype=torch.bfloat16):
training_batch = self.train_one_step(training_batch)
total_loss = training_batch.total_loss
generator_loss = training_batch.generator_loss
fake_score_loss = training_batch.fake_score_loss
grad_norm = training_batch.grad_norm
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"total_loss": f"{total_loss:.4f}",
"generator_loss": f"{generator_loss:.4f}",
"fake_score_loss": f"{fake_score_loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
"ema": "✓" if (self.generator_ema is not None and self.is_ema_ready()) else "✗",
})
progress_bar.update(1)
if self.global_rank == 0:
log_data = {
"train_total_loss": total_loss,
"train_fake_score_loss": fake_score_loss,
"learning_rate": self.lr_scheduler.get_last_lr()[0],
"fake_score_learning_rate": self.fake_score_lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
}
if (step % self.dfake_gen_update_ratio == 0):
log_data["train_generator_loss"] = generator_loss
if use_vsa:
log_data["VSA_train_sparsity"] = current_vsa_sparsity
if self.generator_ema is not None:
log_data["ema_enabled"] = True
log_data["ema_decay"] = self.training_args.ema_decay
else:
log_data["ema_enabled"] = False
ema_stats = self.get_ema_stats()
log_data.update(ema_stats)
if training_batch.dmd_latent_vis_dict:
dmd_additional_logs = {
"generator_timestep": training_batch.dmd_latent_vis_dict["generator_timestep"].item(),
"dmd_timestep": training_batch.dmd_latent_vis_dict["dmd_timestep"].item(),
}
log_data.update(dmd_additional_logs)
faker_score_additional_logs = {
"fake_score_timestep": training_batch.fake_score_latent_vis_dict["fake_score_timestep"].item(),
}
log_data.update(faker_score_additional_logs)
wandb.log(log_data, step=step)
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
if self.training_args.log_visualization:
self.visualize_intermediate_latents(training_batch, self.training_args, step)
if (self.training_args.training_state_checkpointing_steps > 0
and step % self.training_args.training_state_checkpointing_steps == 0):
print("rank", self.global_rank,
"save training state checkpoint at step", step)
save_distillation_checkpoint(
self.transformer, self.fake_score_transformer,
self.global_rank, self.training_args.output_dir, step,
self.optimizer, self.fake_score_optimizer,
self.train_dataloader, self.lr_scheduler,
self.fake_score_lr_scheduler, self.noise_random_generator,
self.generator_ema)
if self.transformer:
self.transformer.train()
self.sp_group.barrier()
if (self.training_args.weight_only_checkpointing_steps > 0
and step % self.training_args.weight_only_checkpointing_steps == 0):
print("rank", self.global_rank,
"save weight-only checkpoint at step", step)
save_distillation_checkpoint(self.transformer,
self.fake_score_transformer,
self.global_rank,
self.training_args.output_dir,
f"{step}_weight_only",
only_save_generator_weight=True,
generator_ema=self.generator_ema)
if self.training_args.use_ema and self.is_ema_ready():
self.save_ema_weights(self.training_args.output_dir, step)
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
self._log_validation(self.transformer, self.training_args, step)
wandb.finish()
print("rank", self.global_rank,
"save final training state checkpoint at step",
self.training_args.max_train_steps)
save_distillation_checkpoint(
self.transformer, self.fake_score_transformer, self.global_rank,
self.training_args.output_dir, self.training_args.max_train_steps,
self.optimizer, self.fake_score_optimizer, self.train_dataloader,
self.lr_scheduler, self.fake_score_lr_scheduler,
self.noise_random_generator, self.generator_ema)
if self.training_args.use_ema and self.is_ema_ready():
self.save_ema_weights(self.training_args.output_dir, self.training_args.max_train_steps)
if get_sp_group():
cleanup_dist_env_and_memory()
+52 -119
View File
@@ -1,5 +1,4 @@
# SPDX-License-Identifier: Apache-2.0
import gc
import dataclasses
import math
import os
@@ -23,6 +22,7 @@ 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,22 +39,20 @@ 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, get_scheduler, get_sigmas,
load_checkpoint, normalize_dit_input, save_checkpoint,
compute_density_for_timestep_sampling, count_trainable, get_scheduler,
get_sigmas, load_checkpoint, normalize_dit_input, save_checkpoint,
shard_latents_across_sp)
from fastvideo.utils import is_vsa_available, set_random_seed, shallow_asdict
from fastvideo.utils import (is_vmoba_available, 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.
@@ -65,7 +63,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
train_dataloader: StatefulDataLoader
train_loader_iter: Iterator[dict[str, Any]]
current_epoch: int = 0
train_transformer_2: bool = False
def __init__(
self,
@@ -101,7 +98,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.sp_world_size = self.sp_group.world_size
self.local_rank = world_group.local_rank
self.transformer = self.get_module("transformer")
self.transformer_2 = self.get_module("transformer_2", None)
self.seed = training_args.seed
self.set_schemas()
@@ -114,26 +110,17 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
if self.transformer_2 is not None:
self.transformer_2 = apply_activation_checkpointing(
self.transformer_2,
checkpointing_type=training_args.
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(
filter(lambda p: p.requires_grad, params_to_optimize))
# Parse betas from string format "beta1,beta2"
betas_str = training_args.betas
betas = tuple(float(x.strip()) for x in betas_str.split(","))
self.optimizer = torch.optim.AdamW(
params_to_optimize,
lr=training_args.learning_rate,
betas=betas,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
@@ -151,30 +138,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
min_lr_ratio=training_args.min_lr_ratio,
last_epoch=self.init_steps - 1,
)
if self.transformer_2 is not None:
# Ensure transformer_2 has trainable parameters before creating optimizer
self.transformer_2.train()
self.transformer_2.requires_grad_(True)
params_to_optimize_2 = self.transformer_2.parameters()
params_to_optimize_2 = list(
filter(lambda p: p.requires_grad, params_to_optimize_2))
self.optimizer_2 = torch.optim.AdamW(
params_to_optimize_2,
lr=training_args.learning_rate,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
self.lr_scheduler_2 = get_scheduler(
training_args.lr_scheduler,
optimizer=self.optimizer_2,
num_warmup_steps=training_args.lr_warmup_steps,
num_training_steps=training_args.max_train_steps,
num_cycles=training_args.lr_num_cycles,
power=training_args.lr_power,
min_lr_ratio=training_args.min_lr_ratio,
last_epoch=self.init_steps - 1,
)
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
training_args.data_path,
@@ -189,7 +152,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
seed=self.seed)
self.noise_scheduler = noise_scheduler
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
self.num_update_steps_per_epoch = math.ceil(
len(self.train_dataloader) /
training_args.gradient_accumulation_steps * training_args.sp_size /
@@ -215,9 +178,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
def _prepare_training(self, training_batch: TrainingBatch) -> TrainingBatch:
self.transformer.train()
self.optimizer.zero_grad()
if self.transformer_2 is not None:
self.transformer_2.train()
self.optimizer_2.zero_grad()
training_batch.total_loss = 0.0
return training_batch
@@ -264,8 +224,17 @@ class TrainingPipeline(LoRAPipeline, ABC):
generator=self.noise_gen_cuda,
device=latents.device,
dtype=latents.dtype)
timesteps = self._sample_timesteps(batch_size, latents.device)
u = compute_density_for_timestep_sampling(
weighting_scheme=self.training_args.weighting_scheme,
batch_size=batch_size,
generator=self.noise_random_generator,
logit_mean=self.training_args.logit_mean,
logit_std=self.training_args.logit_std,
mode_scale=self.training_args.mode_scale,
)
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
timesteps = self.noise_scheduler.timesteps[indices].to(
device=latents.device)
if self.training_args.sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
sp_group = get_sp_group()
@@ -288,38 +257,6 @@ class TrainingPipeline(LoRAPipeline, ABC):
return training_batch
def _sample_timesteps(self, batch_size, device):
# Determine which model to train based on the boundary timestep
if (self.transformer_2 is not None and self.boundary_timestep is not None and
torch.rand(1, generator=self.noise_random_generator).item() <= self.training_args.boundary_ratio):
self.train_transformer_2 = True
else:
self.train_transformer_2 = False
# Broadcast the decision to all processes
decision = torch.tensor(1.0 if self.train_transformer_2 else 0.0, device=self.device)
dist.broadcast(decision, src=0)
self.train_transformer_2 = decision.item() == 1.0
# Sample u from the appropriate range
u = compute_density_for_timestep_sampling(
weighting_scheme=self.training_args.weighting_scheme,
batch_size=batch_size,
generator=self.noise_random_generator,
logit_mean=self.training_args.logit_mean,
logit_std=self.training_args.logit_std,
mode_scale=self.training_args.mode_scale,
)
boundary_ratio = self.training_args.boundary_ratio
if self.train_transformer_2:
u = (1 - boundary_ratio) + u * boundary_ratio # min: 1 - boundary_ratio, max: 1
else:
u = u * (1 - boundary_ratio) # min: 0, max: 1 - boundary_ratio
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
return self.noise_scheduler.timesteps[indices].to(device=device)
def _build_attention_metadata(
self, training_batch: TrainingBatch) -> TrainingBatch:
latents_shape = training_batch.raw_latent_shape
@@ -335,6 +272,20 @@ 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
@@ -359,7 +310,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":
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN" or vmoba_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VMOBA_ATTN":
assert training_batch.attn_metadata is not None
else:
assert training_batch.attn_metadata is None
@@ -370,12 +321,11 @@ class TrainingPipeline(LoRAPipeline, ABC):
# [1000.0],
# device=training_batch.noisy_model_input.device,
# dtype=torch.bfloat16)
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
with set_forward_context(
current_timestep=training_batch.current_timestep,
attn_metadata=training_batch.attn_metadata):
model_pred = current_model(**input_kwargs)
model_pred = self.transformer(**input_kwargs)
if self.training_args.precondition_outputs:
assert training_batch.sigmas is not None
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
@@ -406,12 +356,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
# the following:
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
if max_grad_norm is not None:
# Only clip gradients for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
model_parts = [self.transformer_2]
else:
model_parts = [self.transformer]
model_parts = [self.transformer]
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for m in model_parts for p in m.parameters()],
max_grad_norm,
@@ -456,14 +401,9 @@ class TrainingPipeline(LoRAPipeline, ABC):
training_batch = self._clip_grad_norm(training_batch)
# Only step the optimizer and scheduler for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
self.optimizer_2.step()
self.lr_scheduler_2.step()
else:
self.optimizer.step()
self.lr_scheduler.step()
self.optimizer.step()
self.lr_scheduler.step()
training_batch.total_loss = training_batch.total_loss
training_batch.grad_norm = training_batch.grad_norm
return training_batch
@@ -491,15 +431,10 @@ class TrainingPipeline(LoRAPipeline, ABC):
local_main_process_only=False)
if not self.post_init_called:
self.post_init()
num_trainable_params = _get_trainable_params(self.transformer)
num_trainable_params = count_trainable(self.transformer)
logger.info("Starting training with %s B trainable parameters",
round(num_trainable_params / 1e9, 3))
if getattr(self, "transformer_2", None) is not None:
num_trainable_params = _get_trainable_params(self.transformer_2)
logger.info("Transformer 2: Starting training with %s B trainable parameters",
round(num_trainable_params / 1e9, 3))
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
@@ -520,7 +455,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
self._log_training_info()
self._log_validation(self.training_args,
self._log_validation(self.transformer, self.training_args,
self.init_steps)
# Train!
@@ -541,6 +476,9 @@ 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
@@ -582,10 +520,10 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.transformer.train()
self.sp_group.barrier()
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
self._log_validation(self.training_args, step)
self._log_validation(self.transformer, self.training_args, step)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
trainable_params = round(
_get_trainable_params(self.transformer) / 1e9, 3)
count_trainable(self.transformer) / 1e9, 3)
logger.info(
"GPU memory usage after validation: %s MB, trainable params: %sB",
gpu_memory_usage, trainable_params)
@@ -621,7 +559,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(_get_trainable_params(self.transformer) / 1e9, 3))
round(count_trainable(self.transformer) / 1e9, 3))
# print dtype
logger.info(" Master weight dtype: %s",
self.transformer.parameters().__next__().dtype)
@@ -663,12 +601,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
return batch
@torch.no_grad()
def _log_validation(self, training_args, global_step) -> None:
def _log_validation(self, transformer, training_args, global_step) -> None:
"""
Generate a validation video and log it to wandb to check the quality during training.
"""
training_args.inference_mode = True
training_args.dit_cpu_offload = False
training_args.dit_cpu_offload = True
if not training_args.log_validation:
return
if self.validation_pipeline is None:
@@ -689,10 +627,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
validation_dataloader = DataLoader(validation_dataset,
batch_size=None,
num_workers=0)
self.transformer.eval()
if getattr(self, "transformer_2", None) is not None:
self.transformer_2.eval()
transformer.eval()
validation_steps = training_args.validation_sampling_steps.split(",")
validation_steps = [int(step) for step in validation_steps]
@@ -784,6 +719,4 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Re-enable gradients for training
training_args.inference_mode = False
self.transformer.train()
if getattr(self, "transformer_2", None) is not None:
self.transformer_2.train()
transformer.train()
@@ -1,797 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import dataclasses
import math
import os
import time
from abc import ABC, abstractmethod
from collections import deque
from collections.abc import Iterator
from typing import Any
import imageio
import numpy as np
import torch
import torch.distributed as dist
import torchvision
from diffusers import FlowMatchEulerDiscreteScheduler
from einops import rearrange
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm.auto import tqdm
import fastvideo.envs as envs
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadataBuilder)
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import build_parquet_map_style_dataloader
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
from fastvideo.dataset.validation_dataset import ValidationDataset
from fastvideo.distributed import (cleanup_dist_env_and_memory,
get_local_torch_device, get_sp_group,
get_world_group)
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.pipelines import (ComposedPipelineBase, ForwardBatch,
LoRAPipeline, TrainingBatch)
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, get_scheduler, get_sigmas,
load_checkpoint, normalize_dit_input, save_checkpoint,
shard_latents_across_sp)
from fastvideo.utils import is_vsa_available, set_random_seed, shallow_asdict
import wandb # isort: skip
vsa_available = is_vsa_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.
All reusable components and code should be implemented in this class.
"""
_required_config_modules = ["scheduler", "transformer"]
validation_pipeline: ComposedPipelineBase
train_dataloader: StatefulDataLoader
train_loader_iter: Iterator[dict[str, Any]]
current_epoch: int = 0
train_transformer_2: bool = False
def __init__(
self,
model_path: str,
fastvideo_args: TrainingArgs,
required_config_modules: list[str] | None = None,
loaded_modules: dict[str, torch.nn.Module] | None = None) -> None:
fastvideo_args.inference_mode = False
self.lora_training = fastvideo_args.lora_training
if self.lora_training and fastvideo_args.lora_rank is None:
raise ValueError("lora rank must be set when using lora training")
set_random_seed(fastvideo_args.seed) # for lora param init
super().__init__(model_path, fastvideo_args, required_config_modules,
loaded_modules) # type: ignore
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
raise RuntimeError(
"create_pipeline_stages should not be called for training pipeline")
def set_schemas(self) -> None:
self.train_dataset_schema = pyarrow_schema_t2v
def initialize_training_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing training pipeline...")
self.device = get_local_torch_device()
self.training_args = training_args
world_group = get_world_group()
self.world_size = world_group.world_size
self.global_rank = world_group.rank
self.sp_group = get_sp_group()
self.rank_in_sp_group = self.sp_group.rank_in_group
self.sp_world_size = self.sp_group.world_size
self.local_rank = world_group.local_rank
self.transformer = self.get_module("transformer")
self.transformer_2 = self.get_module("transformer_2", None)
self.seed = training_args.seed
self.set_schemas()
# Set random seeds for deterministic training
assert self.seed is not None, "seed must be set"
set_random_seed(self.seed)
self.transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
self.transformer = apply_activation_checkpointing(
self.transformer,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
if self.transformer_2 is not None:
self.transformer_2 = apply_activation_checkpointing(
self.transformer_2,
checkpointing_type=training_args.
enable_gradient_checkpointing_type)
noise_scheduler = self.modules["scheduler"]
self.set_trainable()
params_to_optimize = self.transformer.parameters()
params_to_optimize = list(
filter(lambda p: p.requires_grad, params_to_optimize))
self.optimizer = torch.optim.AdamW(
params_to_optimize,
lr=training_args.learning_rate,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
self.init_steps = 0
logger.info("optimizer: %s", self.optimizer)
self.lr_scheduler = get_scheduler(
training_args.lr_scheduler,
optimizer=self.optimizer,
num_warmup_steps=training_args.lr_warmup_steps,
num_training_steps=training_args.max_train_steps,
num_cycles=training_args.lr_num_cycles,
power=training_args.lr_power,
min_lr_ratio=training_args.min_lr_ratio,
last_epoch=self.init_steps - 1,
)
if self.transformer_2 is not None:
# Ensure transformer_2 has trainable parameters before creating optimizer
self.transformer_2.train()
self.transformer_2.requires_grad_(True)
params_to_optimize_2 = self.transformer_2.parameters()
params_to_optimize_2 = list(
filter(lambda p: p.requires_grad, params_to_optimize_2))
self.optimizer_2 = torch.optim.AdamW(
params_to_optimize_2,
lr=training_args.learning_rate,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
self.lr_scheduler_2 = get_scheduler(
training_args.lr_scheduler,
optimizer=self.optimizer_2,
num_warmup_steps=training_args.lr_warmup_steps,
num_training_steps=training_args.max_train_steps,
num_cycles=training_args.lr_num_cycles,
power=training_args.lr_power,
min_lr_ratio=training_args.min_lr_ratio,
last_epoch=self.init_steps - 1,
)
self.train_dataset, self.train_dataloader = build_parquet_map_style_dataloader(
training_args.data_path,
training_args.train_batch_size,
parquet_schema=self.train_dataset_schema,
num_data_workers=training_args.dataloader_num_workers,
cfg_rate=training_args.training_cfg_rate,
drop_last=True,
text_padding_length=training_args.pipeline_config.
text_encoder_configs[0].arch_config.
text_len, # type: ignore[attr-defined]
seed=self.seed)
self.noise_scheduler = noise_scheduler
self.boundary_timestep = self.training_args.boundary_ratio * self.noise_scheduler.num_train_timesteps
self.num_update_steps_per_epoch = math.ceil(
len(self.train_dataloader) /
training_args.gradient_accumulation_steps * training_args.sp_size /
training_args.train_sp_batch_size)
self.num_train_epochs = math.ceil(training_args.max_train_steps /
self.num_update_steps_per_epoch)
# TODO(will): is there a cleaner way to track epochs?
self.current_epoch = 0
if self.global_rank == 0:
project = training_args.tracker_project_name or "fastvideo"
wandb_config = dataclasses.asdict(training_args)
wandb.init(project=project,
config=wandb_config,
name=training_args.wandb_run_name)
@abstractmethod
def initialize_validation_pipeline(self, training_args: TrainingArgs):
raise NotImplementedError(
"Training pipelines must implement this method")
def _prepare_training(self, training_batch: TrainingBatch) -> TrainingBatch:
# At the beginning, disable training for all models
self._disable_training(self.transformer, self.optimizer)
if self.transformer_2 is not None:
self._disable_training(self.transformer_2, self.optimizer_2)
training_batch.total_loss = 0.0
return training_batch
def _enable_training(self, model: torch.nn.Module, optimizer: torch.optim.Optimizer) -> None:
"""Enable training mode and gradients for the specified model."""
for param in model.parameters():
param.requires_grad = True
model.train()
optimizer.zero_grad()
def _disable_training(self, model: torch.nn.Module, optimizer: torch.optim.Optimizer) -> None:
"""Disable training mode and gradients for the specified model."""
for param in model.parameters():
param.requires_grad = False
optimizer.zero_grad(set_to_none=True)
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
batch = next(self.train_loader_iter)
# latents, encoder_hidden_states, encoder_attention_mask, infos = batch
latents = batch['vae_latent']
latents = latents[:, :, :self.training_args.num_latent_t]
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
infos = batch['info_list']
training_batch.latents = latents.to(get_local_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = encoder_hidden_states.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
get_local_torch_device(), dtype=torch.bfloat16)
training_batch.infos = infos
return training_batch
def _normalize_dit_input(self,
training_batch: TrainingBatch) -> TrainingBatch:
# TODO(will): support other models
training_batch.latents = normalize_dit_input('wan',
training_batch.latents,
self.get_module("vae"))
return training_batch
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
latents = training_batch.latents
batch_size = latents.shape[0]
noise = torch.randn(latents.shape,
generator=self.noise_gen_cuda,
device=latents.device,
dtype=latents.dtype)
timesteps = self._sample_timesteps(batch_size, latents.device)
# Enable training for the model that will be trained next and disable the other
if self.train_transformer_2:
self._enable_training(self.transformer_2, self.optimizer_2)
self._disable_training(self.transformer, self.optimizer)
else:
self._enable_training(self.transformer, self.optimizer)
if self.transformer_2 is not None:
self._disable_training(self.transformer_2, self.optimizer_2)
if self.training_args.sp_size > 1:
# Make sure that the timesteps are the same across all sp processes.
sp_group = get_sp_group()
sp_group.broadcast(timesteps, src=0)
sigmas = get_sigmas(
self.noise_scheduler,
latents.device,
timesteps,
n_dim=latents.ndim,
dtype=latents.dtype,
)
noisy_model_input = (1.0 -
sigmas) * training_batch.latents + sigmas * noise
training_batch.noisy_model_input = noisy_model_input
training_batch.timesteps = timesteps
training_batch.sigmas = sigmas
training_batch.noise = noise
training_batch.raw_latent_shape = training_batch.latents.shape
return training_batch
def _sample_timesteps(self, batch_size, device):
# Determine which model to train based on the boundary timestep
if (self.transformer_2 is not None and self.boundary_timestep is not None and
torch.rand(1, generator=self.noise_random_generator).item() > self.training_args.boundary_ratio):
self.train_transformer_2 = True
else:
self.train_transformer_2 = False
# Broadcast the decision to all processes
decision = torch.tensor(1.0 if self.train_transformer_2 else 0.0, device=self.device)
dist.broadcast(decision, src=0)
self.train_transformer_2 = decision.item() == 1.0
# Sample u from the appropriate range
u = compute_density_for_timestep_sampling(
weighting_scheme=self.training_args.weighting_scheme,
batch_size=batch_size,
generator=self.noise_random_generator,
logit_mean=self.training_args.logit_mean,
logit_std=self.training_args.logit_std,
mode_scale=self.training_args.mode_scale,
)
boundary_ratio = self.training_args.boundary_ratio
if self.train_transformer_2:
u = boundary_ratio + u * (1.0 - boundary_ratio)
else:
u = u * boundary_ratio
indices = (u * self.noise_scheduler.config.num_train_timesteps).long()
return self.noise_scheduler.timesteps[indices].to(device=device)
def _build_attention_metadata(
self, training_batch: TrainingBatch) -> TrainingBatch:
latents_shape = training_batch.raw_latent_shape
patch_size = self.training_args.pipeline_config.dit_config.patch_size
current_vsa_sparsity = training_batch.current_vsa_sparsity
assert latents_shape is not None
assert training_batch.timesteps is not None
if vsa_available and envs.FASTVIDEO_ATTENTION_BACKEND == "VIDEO_SPARSE_ATTN":
training_batch.attn_metadata = VideoSparseAttentionMetadataBuilder( # type: ignore
).build( # type: ignore
raw_latent_shape=latents_shape[2:5],
current_timestep=training_batch.timesteps,
patch_size=patch_size,
VSA_sparsity=current_vsa_sparsity,
device=get_local_torch_device())
else:
training_batch.attn_metadata = None
return training_batch
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
training_batch.input_kwargs = {
"hidden_states":
training_batch.noisy_model_input,
"encoder_hidden_states":
training_batch.encoder_hidden_states,
"timestep":
training_batch.timesteps.to(get_local_torch_device(),
dtype=torch.bfloat16),
"encoder_attention_mask":
training_batch.encoder_attention_mask,
"return_dict":
False,
}
return training_batch
def _transformer_forward_and_compute_loss(
self, training_batch: TrainingBatch) -> TrainingBatch:
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
input_kwargs = training_batch.input_kwargs
# if 'hunyuan' in self.training_args.model_type:
# input_kwargs["guidance"] = torch.tensor(
# [1000.0],
# device=training_batch.noisy_model_input.device,
# dtype=torch.bfloat16)
current_model = self.transformer_2 if self.train_transformer_2 else self.transformer
with set_forward_context(
current_timestep=training_batch.current_timestep,
attn_metadata=training_batch.attn_metadata):
model_pred = current_model(**input_kwargs)
if self.training_args.precondition_outputs:
assert training_batch.sigmas is not None
model_pred = training_batch.noisy_model_input - model_pred * training_batch.sigmas
assert training_batch.latents is not None
assert training_batch.noise is not None
target = training_batch.latents if self.training_args.precondition_outputs else training_batch.noise - training_batch.latents
# make sure no implicit broadcasting happens
assert model_pred.shape == target.shape, f"model_pred.shape: {model_pred.shape}, target.shape: {target.shape}"
loss = (torch.mean((model_pred.float() - target.float())**2) /
self.training_args.gradient_accumulation_steps)
loss.backward()
avg_loss = loss.detach().clone()
# logger.info(f"rank: {self.rank}, avg_loss: {avg_loss.item()}",
# local_main_process_only=False)
world_group = get_world_group()
world_group.all_reduce(avg_loss, op=dist.ReduceOp.AVG)
training_batch.total_loss += avg_loss.item()
return training_batch
def _clip_grad_norm(self, training_batch: TrainingBatch) -> TrainingBatch:
max_grad_norm = self.training_args.max_grad_norm
# TODO(will): perhaps move this into transformer api so that we can do
# the following:
# grad_norm = transformer.clip_grad_norm_(max_grad_norm)
if max_grad_norm is not None:
# Only clip gradients for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
model_parts = [self.transformer_2]
else:
model_parts = [self.transformer]
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for m in model_parts for p in m.parameters()],
max_grad_norm,
foreach=None,
)
assert grad_norm is not float('nan') or grad_norm is not float(
'inf')
grad_norm = grad_norm.item() if grad_norm is not None else 0.0
else:
grad_norm = 0.0
training_batch.grad_norm = grad_norm
return training_batch
def train_one_step(self, training_batch: TrainingBatch) -> TrainingBatch:
training_batch = self._prepare_training(training_batch)
for _ in range(self.training_args.gradient_accumulation_steps):
training_batch = self._get_next_batch(training_batch)
# Normalize DIT input
training_batch = self._normalize_dit_input(training_batch)
# Create noisy model input
training_batch = self._prepare_dit_inputs(training_batch)
# Shard latents across sp groups
training_batch.latents = shard_latents_across_sp(
training_batch.latents,
num_latent_t=self.training_args.num_latent_t)
# shard noisy_model_input to match
training_batch.noisy_model_input = shard_latents_across_sp(
training_batch.noisy_model_input,
num_latent_t=self.training_args.num_latent_t)
# shard noise to match latents
training_batch.noise = shard_latents_across_sp(
training_batch.noise,
num_latent_t=self.training_args.num_latent_t)
training_batch = self._build_attention_metadata(training_batch)
training_batch = self._build_input_kwargs(training_batch)
training_batch = self._transformer_forward_and_compute_loss(
training_batch)
training_batch = self._clip_grad_norm(training_batch)
# Only step the optimizer and scheduler for the model that is currently training
if self.train_transformer_2 and self.transformer_2 is not None:
self.optimizer_2.step()
self.lr_scheduler_2.step()
else:
self.optimizer.step()
self.lr_scheduler.step()
training_batch.total_loss = training_batch.total_loss
training_batch.grad_norm = training_batch.grad_norm
return training_batch
def _resume_from_checkpoint(self) -> None:
logger.info("Loading checkpoint from %s",
self.training_args.resume_from_checkpoint)
resumed_step = load_checkpoint(
self.transformer, self.global_rank,
self.training_args.resume_from_checkpoint, self.optimizer,
self.train_dataloader, self.lr_scheduler,
self.noise_random_generator)
if resumed_step > 0:
self.init_steps = resumed_step
logger.info("Successfully resumed from step %s", resumed_step)
else:
logger.warning("Failed to load checkpoint, starting from step 0")
self.init_steps = 0
def train(self) -> None:
assert self.seed is not None, "seed must be set"
set_random_seed(self.seed + self.global_rank)
logger.info('rank: %s: start training',
self.global_rank,
local_main_process_only=False)
if not self.post_init_called:
self.post_init()
num_trainable_params = _get_trainable_params(self.transformer)
logger.info("Starting training with %s B trainable parameters",
round(num_trainable_params / 1e9, 3))
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
self.noise_gen_cuda = torch.Generator(device="cuda").manual_seed(
self.seed)
self.validation_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed)
logger.info("Initialized random seeds with seed: %s", self.seed)
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
if self.training_args.resume_from_checkpoint:
self._resume_from_checkpoint()
self.train_loader_iter = iter(self.train_dataloader)
step_times: deque[float] = deque(maxlen=100)
self._log_training_info()
self._log_validation(self.transformer, self.training_args,
self.init_steps)
# Train!
progress_bar = tqdm(
range(0, self.training_args.max_train_steps),
initial=self.init_steps,
desc="Steps",
# Only show the progress bar once on each machine.
disable=self.local_rank > 0,
)
for step in range(self.init_steps + 1,
self.training_args.max_train_steps + 1):
start_time = time.perf_counter()
if vsa_available:
vsa_sparsity = self.training_args.VSA_sparsity
vsa_decay_rate = self.training_args.VSA_decay_rate
vsa_decay_interval_steps = self.training_args.VSA_decay_interval_steps
current_decay_times = min(step // vsa_decay_interval_steps,
vsa_sparsity // vsa_decay_rate)
current_vsa_sparsity = current_decay_times * vsa_decay_rate
else:
current_vsa_sparsity = 0.0
training_batch = TrainingBatch()
training_batch.current_timestep = step
training_batch.current_vsa_sparsity = current_vsa_sparsity
training_batch = self.train_one_step(training_batch)
loss = training_batch.total_loss
grad_norm = training_batch.grad_norm
step_time = time.perf_counter() - start_time
step_times.append(step_time)
avg_step_time = sum(step_times) / len(step_times)
progress_bar.set_postfix({
"loss": f"{loss:.4f}",
"step_time": f"{step_time:.2f}s",
"grad_norm": grad_norm,
})
progress_bar.update(1)
if self.global_rank == 0:
wandb.log(
{
"train_loss": loss,
"learning_rate": self.lr_scheduler.get_last_lr()[0],
"step_time": step_time,
"avg_step_time": avg_step_time,
"grad_norm": grad_norm,
"vsa_sparsity": current_vsa_sparsity,
},
step=step,
)
if step % self.training_args.checkpointing_steps == 0:
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir, step,
self.optimizer, self.train_dataloader,
self.lr_scheduler, self.noise_random_generator)
self.transformer.train()
self.sp_group.barrier()
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
self._log_validation(self.transformer, self.training_args, step)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
trainable_params = round(
_get_trainable_params(self.transformer) / 1e9, 3)
logger.info(
"GPU memory usage after validation: %s MB, trainable params: %sB",
gpu_memory_usage, trainable_params)
wandb.finish()
save_checkpoint(self.transformer, self.global_rank,
self.training_args.output_dir,
self.training_args.max_train_steps, self.optimizer,
self.train_dataloader, self.lr_scheduler,
self.noise_random_generator)
if get_sp_group():
cleanup_dist_env_and_memory()
def _log_training_info(self) -> None:
total_batch_size = (self.world_size *
self.training_args.gradient_accumulation_steps /
self.training_args.sp_size *
self.training_args.train_sp_batch_size)
logger.info("***** Running training *****")
logger.info(" Num examples = %s", len(self.train_dataset))
logger.info(" Dataloader size = %s", len(self.train_dataloader))
logger.info(" Num Epochs = %s", self.num_train_epochs)
logger.info(" Resume training from step %s",
self.init_steps) # type: ignore
logger.info(" Instantaneous batch size per device = %s",
self.training_args.train_batch_size)
logger.info(
" Total train batch size (w. data & sequence parallel, accumulation) = %s",
total_batch_size)
logger.info(" Gradient Accumulation steps = %s",
self.training_args.gradient_accumulation_steps)
logger.info(" Total optimization steps = %s",
self.training_args.max_train_steps)
logger.info(" Total training parameters per FSDP shard = %s B",
round(_get_trainable_params(self.transformer) / 1e9, 3))
# print dtype
logger.info(" Master weight dtype: %s",
self.transformer.parameters().__next__().dtype)
gpu_memory_usage = torch.cuda.memory_allocated() / 1024**2
logger.info("GPU memory usage before train_one_step: %s MB",
gpu_memory_usage)
logger.info("VSA validation sparsity: %s",
self.training_args.VSA_sparsity)
def _prepare_validation_batch(self, sampling_param: SamplingParam,
training_args: TrainingArgs,
validation_batch: dict[str, Any],
num_inference_steps: int) -> ForwardBatch:
sampling_param.prompt = validation_batch['prompt']
sampling_param.height = training_args.num_height
sampling_param.width = training_args.num_width
sampling_param.num_inference_steps = num_inference_steps
sampling_param.data_type = "video"
assert self.seed is not None
sampling_param.seed = self.seed
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
sampling_param.num_frames = num_frames
batch = ForwardBatch(
**shallow_asdict(sampling_param),
latents=None,
generator=self.validation_random_generator,
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
return batch
@torch.no_grad()
def _log_validation(self, transformer, training_args, global_step) -> None:
"""
Generate a validation video and log it to wandb to check the quality during training.
"""
training_args.inference_mode = True
training_args.dit_cpu_offload = True
if not training_args.log_validation:
return
if self.validation_pipeline is None:
raise ValueError("Validation pipeline is not set")
logger.info("Starting validation")
# Create sampling parameters if not provided
sampling_param = SamplingParam.from_pretrained(training_args.model_path)
# Prepare validation prompts
logger.info('rank: %s: fastvideo_args.validation_dataset_file: %s',
self.global_rank,
training_args.validation_dataset_file,
local_main_process_only=False)
validation_dataset = ValidationDataset(
training_args.validation_dataset_file)
validation_dataloader = DataLoader(validation_dataset,
batch_size=None,
num_workers=0)
transformer.eval()
validation_steps = training_args.validation_sampling_steps.split(",")
validation_steps = [int(step) for step in validation_steps]
validation_steps = [step for step in validation_steps if step > 0]
# Log validation results for this step
world_group = get_world_group()
num_sp_groups = world_group.world_size // self.sp_group.world_size
# Process each validation prompt for each validation step
for num_inference_steps in validation_steps:
logger.info("rank: %s: num_inference_steps: %s",
self.global_rank,
num_inference_steps,
local_main_process_only=False)
step_videos: list[np.ndarray] = []
step_captions: list[str] = []
for validation_batch in validation_dataloader:
batch = self._prepare_validation_batch(sampling_param,
training_args,
validation_batch,
num_inference_steps)
logger.info("rank: %s: rank_in_sp_group: %s, batch.prompt: %s",
self.global_rank,
self.rank_in_sp_group,
batch.prompt,
local_main_process_only=False)
assert batch.prompt is not None and isinstance(
batch.prompt, str)
step_captions.append(batch.prompt)
# Run validation inference
output_batch = self.validation_pipeline.forward(
batch, training_args)
samples = output_batch.output
if self.rank_in_sp_group != 0:
continue
# Process outputs
video = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in video:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
step_videos.append(frames)
# Only sp_group leaders (rank_in_sp_group == 0) need to send their
# results to global rank 0
if self.rank_in_sp_group == 0:
if self.global_rank == 0:
# Global rank 0 collects results from all sp_group leaders
all_videos = step_videos # Start with own results
all_captions = step_captions
# Receive from other sp_group leaders
for sp_group_idx in range(1, num_sp_groups):
src_rank = sp_group_idx * self.sp_world_size # Global rank of other sp_group leaders
recv_videos = world_group.recv_object(src=src_rank)
recv_captions = world_group.recv_object(src=src_rank)
all_videos.extend(recv_videos)
all_captions.extend(recv_captions)
video_filenames = []
for i, (video, caption) in enumerate(
zip(all_videos, all_captions, strict=True)):
os.makedirs(training_args.output_dir, exist_ok=True)
filename = os.path.join(
training_args.output_dir,
f"validation_step_{global_step}_inference_steps_{num_inference_steps}_video_{i}.mp4"
)
imageio.mimsave(filename, video, fps=sampling_param.fps)
video_filenames.append(filename)
logs = {
f"validation_videos_{num_inference_steps}_steps": [
wandb.Video(filename, caption=caption)
for filename, caption in zip(
video_filenames, all_captions, strict=True)
]
}
wandb.log(logs, step=global_step)
else:
# Other sp_group leaders send their results to global rank 0
world_group.send_object(step_videos, dst=0)
world_group.send_object(step_captions, dst=0)
# Re-enable gradients for training
training_args.inference_mode = False
transformer.train()
+3 -168
View File
@@ -202,7 +202,6 @@ def save_distillation_checkpoint(generator_transformer,
generator_scheduler=None,
fake_score_scheduler=None,
noise_generator=None,
generator_ema=None,
only_save_generator_weight=False) -> None:
"""
Save distillation checkpoint with both generator and fake_score models.
@@ -234,8 +233,6 @@ def save_distillation_checkpoint(generator_transformer,
if generator_scheduler is not None:
generator_states["scheduler"] = SchedulerWrapper(
generator_scheduler)
if generator_ema is not None:
generator_states["ema"] = generator_ema.state_dict()
generator_dcp_dir = os.path.join(save_dir, "distributed_checkpoint",
"generator")
@@ -405,8 +402,7 @@ def load_distillation_checkpoint(generator_transformer,
dataloader=None,
generator_scheduler=None,
fake_score_scheduler=None,
noise_generator=None,
generator_ema=None) -> int:
noise_generator=None) -> int:
"""
Load distillation checkpoint with both generator and fake_score models.
Returns the step number from which training should resume.
@@ -460,18 +456,6 @@ def load_distillation_checkpoint(generator_transformer,
end_time - begin_time,
local_main_process_only=False)
# Load EMA state if available and generator_ema is provided
if generator_ema is not None:
try:
ema_state = generator_states.get("ema")
if ema_state is not None:
generator_ema.load_state_dict(ema_state)
logger.info("rank: %s, generator EMA state loaded successfully", rank)
else:
logger.info("rank: %s, no EMA state found in checkpoint", rank)
except Exception as e:
logger.warning("rank: %s, failed to load EMA state: %s", rank, str(e))
# Load critic distributed checkpoint
critic_dcp_dir = os.path.join(checkpoint_path, "distributed_checkpoint",
"critic")
@@ -1296,154 +1280,5 @@ def get_scheduler(
last_epoch=last_epoch)
class EMA_FSDP:
"""
FSDP2-friendly EMA with two modes:
- mode="local_shard" (default): maintain float32 CPU EMA of local parameter shards on every rank.
Provides a context manager to temporarily swap EMA weights into the live model for teacher forward.
- mode="rank0_full": maintain a consolidated float32 CPU EMA of full parameters on rank 0 only
using gather_state_dict_on_cpu_rank0(). Useful for checkpoint export; not for teacher forward.
Usage (local_shard for CM teacher):
ema = EMA_FSDP(model, decay=0.999, mode="local_shard")
for step in ...:
ema.update(model)
with ema.apply_to_model(model):
with torch.no_grad():
y_teacher = model(...)
Usage (rank0_full for export):
ema = EMA_FSDP(model, decay=0.999, mode="rank0_full")
ema.update(model)
ema.state_dict() # on rank 0
"""
def __init__(self, module, decay: float = 0.999, mode: str = "local_shard"):
self.decay = float(decay)
self.mode = mode
self.shadow: dict[str, torch.Tensor] = {}
self.rank = dist.get_rank() if dist.is_initialized() else 0
if self.mode not in {"local_shard", "rank0_full"}:
raise ValueError(f"Unsupported EMA_FSDP mode: {self.mode}")
self._init_shadow(module)
@staticmethod
def _to_local_tensor(t: torch.Tensor) -> torch.Tensor:
# DTensor-aware to_local fetch; fall back to raw tensor
try:
from torch.distributed.tensor import DTensor # type: ignore
if isinstance(t, DTensor):
return t.to_local()
except Exception:
pass
return t
@torch.no_grad()
def _init_shadow(self, module):
if self.mode == "rank0_full":
cpu_state = gather_state_dict_on_cpu_rank0(module, device=None)
if self.rank == 0:
self.shadow = {k: v.detach().clone().float().cpu() for k, v in cpu_state.items()}
else:
self.shadow = {}
return
# local_shard: maintain EMA of local shards for requires_grad params
self.shadow = {}
for name, p in module.named_parameters():
if not p.requires_grad:
continue
local = self._to_local_tensor(p.detach())
self.shadow[name] = local.clone().float().cpu()
@torch.no_grad()
def update(self, module):
d = self.decay
if self.mode == "rank0_full":
if self.rank != 0:
return
cpu_state = gather_state_dict_on_cpu_rank0(module, device=None)
for n, v in cpu_state.items():
v_cpu = v.detach().float().cpu()
if n not in self.shadow:
self.shadow[n] = v_cpu.clone()
else:
self.shadow[n].mul_(d).add_(v_cpu, alpha=1.0 - d)
return
# local_shard: update local shard EMA on every rank
for name, p in module.named_parameters():
if not p.requires_grad:
continue
local = self._to_local_tensor(p.detach())
v_cpu = local.float().cpu()
if name not in self.shadow:
self.shadow[name] = v_cpu.clone()
else:
self.shadow[name].mul_(d).add_(v_cpu, alpha=1.0 - d)
def state_dict(self) -> dict[str, torch.Tensor]:
if self.mode == "rank0_full":
return {k: v.clone() for k, v in self.shadow.items()} if self.rank == 0 else {}
return {k: v.clone() for k, v in self.shadow.items()}
def load_state_dict(self, sd: dict[str, torch.Tensor]):
self.shadow = {k: v.clone() for k, v in sd.items()}
@torch.no_grad()
def copy_to_unwrapped(self, module) -> None:
"""
Copy EMA weights into a non-sharded (unwrapped) module. Intended for export/eval.
For mode="rank0_full", only rank 0 has the full EMA state.
"""
if self.mode == "rank0_full" and self.rank != 0:
return
name_to_param = dict(module.named_parameters())
for n, w in self.shadow.items():
if n in name_to_param:
p = name_to_param[n]
p.data.copy_(w.to(dtype=p.dtype, device=p.device))
class _ApplyEMACtx:
def __init__(self, ema: "EMA_FSDP", module):
self.ema = ema
self.module = module
self.saved: dict[str, torch.Tensor] = {}
def __enter__(self):
if self.ema.mode != "local_shard":
raise RuntimeError("EMA apply_to_model is only supported for mode='local_shard'")
with torch.no_grad():
for name, p in self.module.named_parameters():
if not p.requires_grad:
continue
# Save local shard
p_local = EMA_FSDP._to_local_tensor(p.detach())
if p_local.numel() == 0:
# Nothing to swap on this rank for this param
continue
self.saved[name] = p_local.clone().to(device=p_local.device, dtype=p_local.dtype)
if name in self.ema.shadow:
ema_cpu = self.ema.shadow[name]
if ema_cpu.numel() != p_local.numel():
# Shard shape mismatch (e.g., empty shard here), skip
continue
# Copy EMA shard into local param shard
p_local.copy_(ema_cpu.to(dtype=p_local.dtype, device=p_local.device))
return self.module
def __exit__(self, exc_type, exc, tb):
with torch.no_grad():
for name, p in self.module.named_parameters():
if name in self.saved:
p_local = EMA_FSDP._to_local_tensor(p.detach())
if p_local.numel() == 0:
continue
saved_local = self.saved[name]
if saved_local.numel() != p_local.numel():
continue
p_local.copy_(saved_local)
self.saved.clear()
return False
def apply_to_model(self, module):
return EMA_FSDP._ApplyEMACtx(self, module)
def count_trainable(model: torch.nn.Module) -> int:
return sum(p.numel() for p in model.parameters() if p.requires_grad)
@@ -1,77 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler)
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import WanCausalDMDPipeline
from fastvideo.training.self_forcing_distillation_pipeline import SelfForcingDistillationPipeline
from fastvideo.utils import is_vsa_available
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanSelfForcingDistillationPipeline(SelfForcingDistillationPipeline):
"""
A self-forcing distillation pipeline for Wan that uses the self-forcing methodology
with DMD for video generation.
"""
_required_config_modules = [
"scheduler", "transformer", "vae", "real_score_transformer",
"fake_score_transformer"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
"""Initialize Wan-specific scheduler."""
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
validation_pipeline = WanCausalDMDPipeline.from_pretrained(
training_args.model_path,
args=args_copy, # type: ignore
inference_mode=True,
loaded_modules={"transformer": self.get_module("transformer")},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
pin_cpu_memory=training_args.pin_cpu_memory,
dit_cpu_offload=True)
self.validation_pipeline = validation_pipeline
def main(args) -> None:
logger.info("Starting Wan self-forcing distillation pipeline...")
pipeline = WanSelfForcingDistillationPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("Wan self-forcing distillation pipeline completed")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
main(args)
+2 -3
View File
@@ -34,7 +34,7 @@ class WanTrainingPipeline(TrainingPipeline):
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
assert training_args.dit_cpu_offload == False
args_copy.inference_mode = True
validation_pipeline = WanPipeline.from_pretrained(
training_args.model_path,
@@ -42,13 +42,12 @@ class WanTrainingPipeline(TrainingPipeline):
inference_mode=True,
loaded_modules={
"transformer": self.get_module("transformer"),
"transformer_2": self.get_module("transformer_2")
},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
pin_cpu_memory=training_args.pin_cpu_memory,
dit_cpu_offload=training_args.dit_cpu_offload)
dit_cpu_offload=True)
self.validation_pipeline = validation_pipeline
+30 -1
View File
@@ -23,10 +23,14 @@ 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
@@ -34,7 +38,6 @@ from torch.distributed.fsdp import MixedPrecisionPolicy
import fastvideo.envs as envs
from fastvideo.logger import init_logger
logger = init_logger(__name__)
T = TypeVar("T")
@@ -615,6 +618,7 @@ def maybe_download_model_index(model_name_or_path: str) -> dict[str, Any]:
f"Failed to download or parse model_index.json for {model_name_or_path}: {e}"
) from e
def update_environment_variables(envs: dict[str, str]):
for k, v in envs.items():
if k in os.environ and os.environ[k] != v:
@@ -814,6 +818,17 @@ 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,
@@ -875,3 +890,17 @@ 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")
-204
View File
@@ -1,204 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Encoding stage for diffusion pipelines.
"""
from typing import Optional
import PIL.Image
import torch
from fastvideo.v1.distributed import get_torch_device
from fastvideo.v1.fastvideo_args import FastVideoArgs
from fastvideo.v1.logger import init_logger
from fastvideo.v1.models.vaes.common import ParallelTiledVAE
from fastvideo.v1.models.vision_utils import (get_default_height_width,
normalize, numpy_to_pt,
pil_to_numpy, resize)
from fastvideo.v1.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.v1.pipelines.stages.base import PipelineStage
from fastvideo.v1.pipelines.stages.validators import V # Import validators
from fastvideo.v1.pipelines.stages.validators import VerificationResult
from fastvideo.v1.utils import PRECISION_TO_TYPE
logger = init_logger(__name__)
class EncodingStage(PipelineStage):
"""
Stage for encoding pixel representations into latent space.
This stage handles the encoding of pixel representations into the final
input format (e.g., latents).
"""
def __init__(self, vae: ParallelTiledVAE) -> None:
self.vae: ParallelTiledVAE = vae
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
"""
Encode pixel representations into latent space.
Args:
batch: The current batch information.
fastvideo_args: The inference arguments.
Returns:
The batch with encoded outputs.
"""
self.vae = self.vae.to(get_torch_device())
assert batch.height is not None
assert batch.width is not None
latent_height = batch.height // self.vae.spatial_compression_ratio
latent_width = batch.width // self.vae.spatial_compression_ratio
image = batch.preprocessed_image
# TODO(will)
if image is None:
assert batch.pil_image is not None
image = batch.pil_image
image = self.preprocess(
image,
vae_scale_factor=self.vae.spatial_compression_ratio,
height=batch.height,
width=batch.width).to(get_torch_device(), dtype=torch.float32)
image = image.unsqueeze(2)
else:
# assumes image is loaded from parquet file and used for validation
torch.distributed.breakpoint()
image = image.transpose(1, 2)
logger.info("image: %s", image.shape)
video_condition = torch.cat([
image,
image.new_zeros(image.shape[0], image.shape[1],
batch.num_frames - 1, batch.height, batch.width)
],
dim=2)
video_condition = video_condition.to(device=get_torch_device(),
dtype=torch.float32)
# 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
# Encode Image
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:
video_condition = video_condition.to(vae_dtype)
encoder_output = self.vae.encode(video_condition)
generator = batch.generator
if generator is None:
raise ValueError("Generator must be provided")
latent_condition = self.retrieve_latents(encoder_output, generator)
# 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):
latent_condition -= self.vae.shift_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition -= self.vae.shift_factor
if isinstance(self.vae.scaling_factor, torch.Tensor):
latent_condition = latent_condition * self.vae.scaling_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition = latent_condition * self.vae.scaling_factor
mask_lat_size = torch.ones(1, 1, batch.num_frames, latent_height,
latent_width)
mask_lat_size[:, :, list(range(1, batch.num_frames))] = 0
first_frame_mask = mask_lat_size[:, :, 0:1]
first_frame_mask = torch.repeat_interleave(
first_frame_mask,
dim=2,
repeats=self.vae.temporal_compression_ratio)
mask_lat_size = torch.concat(
[first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
mask_lat_size = mask_lat_size.view(1, -1,
self.vae.temporal_compression_ratio,
latent_height, latent_width)
mask_lat_size = mask_lat_size.transpose(1, 2)
mask_lat_size = mask_lat_size.to(latent_condition.device)
batch.image_latent = torch.concat([mask_lat_size, latent_condition],
dim=1)
# Offload models if needed
if hasattr(self, 'maybe_free_model_hooks'):
self.maybe_free_model_hooks()
return batch
def retrieve_latents(self,
encoder_output: torch.Tensor,
generator: Optional[torch.Generator] = None,
sample_mode: str = "sample"):
if sample_mode == "sample":
return encoder_output.sample(generator)
elif sample_mode == "argmax":
return encoder_output.mode()
else:
raise AttributeError(
"Could not access latents of provided encoder_output")
def preprocess(
self,
image: PIL.Image.Image,
vae_scale_factor: int,
height: Optional[int] = None,
width: Optional[int] = None,
resize_mode: str = "default", # "default", "fill", "crop"
) -> torch.Tensor:
image = [image]
height, width = get_default_height_width(image[0], vae_scale_factor,
height, width)
image = [
resize(i, height, width, resize_mode=resize_mode) for i in image
]
image = pil_to_numpy(image) # to np
image = numpy_to_pt(image) # to pt
do_normalize = True
if image.min() < 0:
do_normalize = False
if do_normalize:
image = normalize(image)
return image
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify encoding stage inputs."""
result = VerificationResult()
# result.add_check("pil_image", batch.pil_image)
result.add_check("height", batch.height, V.positive_int)
result.add_check("width", batch.width, V.positive_int)
result.add_check("generator", batch.generator,
V.generator_or_list_generators)
result.add_check("num_frames", batch.num_frames, V.positive_int)
return result
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify encoding stage outputs."""
result = VerificationResult()
result.add_check("image_latent", batch.image_latent,
[V.is_tensor, V.with_dims(5)])
return result
+37 -159
View File
@@ -1,19 +1,19 @@
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
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,31 +154,17 @@ class VideoForwardBatchBuilder:
class ParquetDatasetSaver:
"""Component for saving and writing Parquet datasets"""
"""Component for saving and writing Parquet datasets using shared parquet_io."""
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
"""
def __init__(self, flush_frequency: int, samples_per_file: int,
schema: pa.Schema,
record_creator: Callable[..., list[dict[str, Any]]]):
self.flush_frequency = flush_frequency
self.samples_per_file = samples_per_file
self.schema_fields = schema_fields
self.schema = schema
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.num_saved_files: int = 0
self._writer: ParquetDatasetWriter | None = None
def save_and_write_parquet_batch(
self,
@@ -226,11 +212,12 @@ class ParquetDatasetSaver:
if batch_data:
self.num_processed_samples += len(batch_data)
# 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)
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)
logger.debug("Collected batch with %s samples", len(table))
# If flush is needed
@@ -260,145 +247,35 @@ class ParquetDatasetSaver:
def _convert_batch_to_pyarrow_table(self,
batch_data: list[dict]) -> pa.Table:
"""Convert batch data to PyArrow table"""
arrays = []
# Deprecated path, kept for backward compatibility if needed.
return records_to_table(batch_data, self.schema)
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]))
def flush_tables(self, output_dir: str, write_remainder: bool = False):
"""Flush buffered records to disk.
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:
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
_ = 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
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 flush_last(self, output_dir: str):
"""Flush and write any remaining rows (final flush)."""
self.flush_tables(output_dir, write_remainder=True)
def clean_up(self) -> None:
"""Clean up all tables"""
self.all_tables = []
self._writer = None
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:
@@ -431,7 +308,8 @@ def build_dataset(preprocess_config: PreprocessConfig, split: str,
return item
dataset = dataset.map(add_video_column)
dataset = dataset.cast_column("video", Video())
if preprocess_config.video_loader_type == VideoLoaderType.TORCHCODEC:
dataset = dataset.cast_column("video", Video())
else:
raise ValueError(
f"Invalid dataset type: {preprocess_config.dataset_type}")
@@ -85,16 +85,16 @@ class PreprocessWorkflow(WorkflowBase):
# record creator
if self.fastvideo_args.workload_type == WorkloadType.I2V:
record_creator = i2v_record_creator
schema_fields = [f.name for f in pyarrow_schema_i2v]
schema = pyarrow_schema_i2v
else:
record_creator = basic_t2v_record_creator
schema_fields = [f.name for f in pyarrow_schema_t2v]
schema = 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_fields=schema_fields,
schema=schema,
record_creator=record_creator,
)
self.add_component("processed_dataset_saver", processed_dataset_saver)
+2 -2
View File
@@ -3,7 +3,7 @@ from huggingface_hub import HfApi
api = HfApi()
api.upload_folder(
folder_path="wow",
repo_id="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
folder_path="Wan2.2-TI2V-5B-Diffusers",
repo_id="FastVideo/FastWan2.2-TI2V-5B-Diffusers",
repo_type="model",
)
+28
View File
@@ -0,0 +1,28 @@
#!/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/
-56
View File
@@ -1,56 +0,0 @@
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export WANDB_API_KEY='50632ebd88ffd970521cec9ab4a1a2d7e85bfc45'
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
export TRITON_CACHE_DIR=/tmp/triton_cache
# DATA_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/Wan-Syn-upload/latents_i2v/train/
DATA_DIR=/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol_processed_t2v/combined_parquet_dataset
# VALIDATION_DIR=/mnt/sharefs/users/hao.zhang/Vchitect-2M/mixkit/validation_8.json
VALIDATION_DIR=/mnt/weka/home/hao.zhang/wei/FastVideo/data/crush-smol-single_processed_t2v/validation.json
NUM_GPUS=8
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
CHECKPOINT_PATH="outputs_train_test/wan_finetune/checkpoint-10"
# 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/wan_training_pipeline.py \
--model_path Wan-AI/Wan2.2-T2V-A14B-Diffusers \
--inference_mode False\
--pretrained_model_name_or_path Wan-AI/Wan2.2-T2V-A14B-Diffusers \
--cache_dir "/home/ray/.cache" \
--data_path "$DATA_DIR" \
--validation_dataset_file "$VALIDATION_DIR" \
--train_batch_size 1 \
--num_latent_t 16\
--sp_size 4 \
--tp_size 1 \
--num_gpus $NUM_GPUS \
--hsdp_replicate_dim 1 \
--hsdp-shard-dim 8 \
--train_sp_batch_size 1 \
--dataloader_num_workers 4 \
--gradient_accumulation_steps 1 \
--max_train_steps 30000 \
--learning_rate 1e-5 \
--mixed_precision "bf16" \
--checkpointing_steps 1000 \
--validation_steps 30 \
--validation_sampling_steps "40" \
--checkpoints_total_limit 3 \
--ema_start_step 0 \
--training_cfg_rate 0.1 \
--seed 1024 \
--output_dir "outputs_train_test/wan_finetune_v1" \
--tracker_project_name VSA_finetune \
--num_height 448 \
--num_width 832 \
--num_frames 61 \
--flow_shift 5 \
--validation_guidance_scale "5.0" \
--num_euler_timesteps 50 \
--master_weight_type "fp32" \
--dit_precision "fp32" \
--weight_decay 0.01 \
--max_grad_norm 1.0