Compare commits

...
Author SHA1 Message Date
JerryZhou54 4ea813265b Resolve timestep mismatch between dit forward and pred_noise_to_pred_video 2025-09-17 03:51:05 +00:00
SolitaryThinker 7e2c8f8d49 comment out vmoba 2025-09-16 04:27:42 +00:00
SolitaryThinker 421dcbb50b Merge branch 'wei/dit_debug' into will/ode_init 2025-09-16 02:14:14 +00:00
SolitaryThinker b9662bf882 lmdb datasets 2025-09-16 02:06:43 +00:00
JerryZhou54 140fb9f20e Fix test for forward_train 2025-09-15 23:23:33 +00:00
JerryZhou54 2078876b98 Ensure 0 numerical diff for forward_train 2025-09-15 22:49:52 +00:00
JerryZhou54 a953f46bd6 Add test for _forward_train 2025-09-15 22:37:50 +00:00
JerryZhou54 adae957008 Fix numerical diff between causal_wanvideo.py and SF's causal wan 2025-09-14 08:09:57 +00:00
SolitaryThinker fa40553afb t2v to i2v finetune
checkpoint ode

checkpoint

fix t2v to i2v

lint

chekpt

checkpoint

hacked but working

ode_init scripts

WIP fixing time embedding

WIP fixing time embedding

checkpoint

update

fix

revert

revert

revert

update

update

visualize
2025-09-14 02:23:42 +00:00
William Lin b93ef4289d [bugfix] Fix empty PipelineConfigs for Wan2.2 A14B (#800) 2025-09-13 17:31:38 -07:00
401bdbd316 [self-forcing] [3/n] Text embed only preprocessing (#797)
Co-authored-by: RandNMR73 <notomatthew31@gmail.com>
Co-authored-by: JerryZhou54 <zhouw.jerry2017@outlook.com>
Co-authored-by: kevin314 <kevin.lin.cs1@gmail.com>
2025-09-13 14:03:53 -07:00
William Lin 1048d79cf8 [bugfix] pin gradio version and set current_vsa_sparsity in TrainingPipeline (#798) 2025-09-11 17:04:47 -07:00
1e8406162d [bugfix] Fix delta calculation (#796)
Co-authored-by: zbchu2 <zbchu2@iflytek.com>
Co-authored-by: William Lin <SolitaryThinker@users.noreply.github.com>
2025-09-11 16:31:23 -07:00
William Lin 03edd35c83 [preprocessing] [self-forcing] [2/n] Improve preprocessing and add ode trajectory dataset schema (#794) 2025-09-10 17:33:57 -07:00
RandNMR73 93ebd15a0d text preprocessing ready 2025-09-10 11:12:48 +00:00
JerryZhou54 1110474065 checkpoint 2025-09-10 08:57:54 +00:00
JerryZhou54 80baffd540 Enable timestep warping & using SelfForcing scheduler 2025-09-09 23:30:55 +00:00
JerryZhou54 918180048e Stop backprop through kv_cache 2025-09-09 10:02:03 +00:00
RandNMR73 b7dbd7cb9e new branch 2025-09-09 10:02:00 +00:00
RandNMR73 71159b6416 inference works after changes added 2025-09-09 10:01:31 +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
144 changed files with 16847 additions and 809 deletions
+34
View File
@@ -176,3 +176,37 @@ 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"
- path:
- "fastvideo/**"
- "pyproject.toml"
- "docker/Dockerfile.python3.12"
config:
command: "timeout 15m .buildkite/scripts/pr_test.sh"
label: "Unit Tests"
env:
- TEST_TYPE=unit_test
agents:
queue: "default"
+13
View File
@@ -109,6 +109,19 @@ 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"
;;
"unit_test")
log "Running unit tests..."
MODAL_COMMAND="$MODAL_ENV python3 -m modal run $MODAL_TEST_FILE::run_unit_test"
;;
*)
log "Error: Unknown test type: $TEST_TYPE"
exit 1
+34 -9
View File
@@ -62,8 +62,8 @@ on:
required: false
default: false
type: boolean
run_nightly_test:
description: "Run nightly-test"
run_unit_test:
description: "Run unit-test"
required: false
default: false
type: boolean
@@ -93,6 +93,7 @@ jobs:
inference-test-STA: ${{ steps.filter.outputs.inference-test-STA }}
precision-test-STA: ${{ steps.filter.outputs.precision-test-STA }}
precision-test-VSA: ${{ steps.filter.outputs.precision-test-VSA }}
unit-test: ${{ steps.filter.outputs.unit-test }}
steps:
- uses: actions/checkout@v4
- uses: dorny/paths-filter@v3
@@ -102,6 +103,8 @@ jobs:
# Define reusable path patterns
common-paths: &common-paths
- 'pyproject.toml'
- 'docker/Dockerfile.python3.10'
- 'docker/Dockerfile.python3.11'
- 'docker/Dockerfile.python3.12'
sta-kernel-paths: &sta-kernel-paths
- 'csrc/attn/sliding_tile_attn/**'
@@ -155,6 +158,9 @@ jobs:
precision-test-VSA:
- *common-paths
- *vsa-kernel-paths
unit-test:
- 'fastvideo/**'
- *common-paths
encoder-test:
needs: change-filter
@@ -333,23 +339,42 @@ jobs:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
nightly-test:
unit-test:
needs: change-filter
if: >-
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
(github.event_name != 'workflow_dispatch' && needs.change-filter.outputs.unit-test == 'true') ||
(github.event_name == 'workflow_dispatch' && github.event.inputs.run_unit_test == 'true')
uses: ./.github/workflows/runpod-test.yml
with:
job_id: "nightly-test"
gpu_type: "NVIDIA A40"
gpu_count: 4
job_id: "unit-test"
gpu_type: "NVIDIA L40S"
gpu_count: 1
volume_size: 100
disk_size: 100
image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
test_command: "uv pip install -e .[test] && pytest ./fastvideo/dataset/ -vs && pytest ./fastvideo/workflow/ -vs"
timeout_minutes: 30
secrets:
RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
# nightly-test:
# if: >-
# (github.event_name == 'workflow_dispatch' && github.event.inputs.run_nightly_test == 'true')
# uses: ./.github/workflows/runpod-test.yml
# with:
# job_id: "nightly-test"
# gpu_type: "NVIDIA A40"
# gpu_count: 4
# volume_size: 100
# disk_size: 100
# image: "ghcr.io/${{ github.repository }}/fastvideo-dev:py3.12-latest"
# test_command: "wandb login $WANDB_API_KEY && uv pip install -e .[test] && pytest ./fastvideo/tests/nightly/test_e2e_overfit_single_sample.py -vs"
# timeout_minutes: 30
# secrets:
# RUNPOD_API_KEY: ${{ secrets.RUNPOD_API_KEY }}
# RUNPOD_PRIVATE_KEY: ${{ secrets.RUNPOD_PRIVATE_KEY }}
# WANDB_API_KEY: ${{ secrets.WANDB_API_KEY }}
runpod-cleanup:
# Add other jobs to this list as you create them
+3
View File
@@ -64,3 +64,6 @@ docs/source/distillation/examples/
!docs/source/_static/images/**/*.png
!comfyui/assets/**/*.png
!comfyui/assets/**/*.gif
dmd_t2v_output/
preprocess_output_text/
+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')
+9
View File
@@ -0,0 +1,9 @@
# VidProm Dataset
From [Self-Forcing](https://github.com/gdhe17/Self-Forcing) repository.
## Download the dataset
```bash
./download_dataset.sh
```
@@ -0,0 +1,3 @@
#! /bin/bash
huggingface-cli download gdhe17/Self-Forcing vidprom_filtered_extended.txt --local-dir prompts
@@ -0,0 +1,3 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -0,0 +1,76 @@
{
"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": "validation_dataset/yYcK4nANZz4-Scene-034.mp4",
"num_inference_steps": 40,
"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": "validation_dataset/yYcK4nANZz4-Scene-027.mp4",
"num_inference_steps": 40,
"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": "validation_dataset/yYcK4nANZz4-Scene-030.mp4",
"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": "validation_dataset/1gGQy4nxyUo-Scene-016.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "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.",
"image_path": null,
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-056.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "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.",
"image_path": null,
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-059.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "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.",
"image_path": null,
"video_path": "validation_dataset/EJqsC21GSBY-Scene-059.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "A colorful puzzle ball is being crushed by a large metal cylinder, which flattens the objects as if they were under a hydraulic press.",
"image_path": null,
"video_path": "validation_dataset/GBSfpTcKegk-Scene-003.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -0,0 +1,13 @@
{
"data": [
{
"caption": "A watermelon wearing a helmet is crushed by a hydraulic press, causing it to flatten and burst open.",
"image_path": null,
"video_path": "validation_dataset/1gGQy4nxyUo-Scene-016.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 77
}
]
}
@@ -0,0 +1,151 @@
#!/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
# Basic Info
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=29501
export TOKENIZERS_PARALLELISM=false
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
# Configs
NUM_GPUS=4
# Model paths for Self-Forcing 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/mixkit-64_processed/Node_0_GPU_1_File_1/combined_parquet_dataset"
VALIDATION_DATASET_FILE="data/mixkit-64_processed/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 # Updated for self-forcing DMD
--output_dir "/mnt/sharefs/users/hao.zhang/SFwan_t2v_finetune"
--max_train_steps 4000
--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 # Must be divisible by num_frame_per_block (81 % 3 = 0 ✓)
--enable_gradient_checkpointing_type "full"
--log_visualization
--simulate_generator_forward
--num_frame_per_block 3 # Frame generation block size for self-forcing
--enable_gradient_masking
--gradient_mask_last_n_frames 21
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS # 64
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1 # 64
--hsdp_shard_dim $NUM_GPUS
)
# 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 100
--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 500
--weight_only_checkpointing_steps 500
--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'
--warp_denoising_step
)
# Self-forcing specific arguments
self_forcing_args=(
--independent_first_frame False # Whether to treat first frame independently
--same_step_across_blocks True # Whether to use same denoising step across all blocks
--last_step_only False # Whether to only use the last denoising step
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
--validate_cache_structure False # Set to True for debugging KV cache issues
)
torchrun \
--nnodes 1 \
--master_port $MASTER_PORT \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}" \
"${dmd_args[@]}" \
"${self_forcing_args[@]}"
@@ -0,0 +1,3 @@
#!/bin/bash
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
@@ -0,0 +1,24 @@
#!/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/merge.txt"
OUTPUT_DIR="data/crush-smol_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 8 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 81 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v"
@@ -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[@]}"
+13 -14
View File
@@ -1,6 +1,6 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples"
def main():
@@ -17,31 +17,30 @@ def main():
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
ti2v_task=True,
# image_encoder_cpu_offload=False,
)
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
# sampling_param.num_frames = 45
# sampling_param.image_path = "https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/astronaut.jpg"
sampling_param.image_path = "test.jpg"
# Generate videos with the same simple API, regardless of GPU count
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
"A girl is packing a suitcase when stuff suddently starts flying around the room."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True)
# video = generator.generate_video(prompt, sampling_param=sampling_param, output_path="wan_t2v_videos/")
# Generate another video with a different prompt, without reloading the
# model!
prompt2 = (
"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.")
video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
# prompt2 = (
# "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.")
# video2 = generator.generate_video(prompt2, output_path=OUTPUT_PATH, save_video=True)
if __name__ == "__main__":
main()
main()
@@ -19,13 +19,15 @@ def main():
)
sampling_param = SamplingParam.from_pretrained(model_name)
sampling_param.num_frames = 81
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
prompts = [
"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.",
"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.",
]
for prompt in prompts:
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
if __name__ == "__main__":
main()
@@ -0,0 +1,99 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
# DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init_5/combined_parquet_dataset/"
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing-extended-t2v-1-3b/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_70k"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "fixed_wan_ode_init_70k_6e-6"
# --resume_from_checkpoint "ode_init_diffusers/"
--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
--enable_gradient_checkpointing_type "full"
)
# 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[@]}"
@@ -0,0 +1,135 @@
#!/bin/bash
#SBATCH --job-name=1e5B2_16kFV_warp_ode_vidprom
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=ode_vidprom16k_warp/Dode_vidprom8b16k_1e-5.out
#SBATCH --error=ode_vidprom16k_warp/Dode_vidprom8b16k_1e-5.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv2
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 WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing-16k-t2v-1-3b-81/"
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/ode_init/validation.json"
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "Dwarp_vidprom_8b16k_test_warp_1e-5"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "Dwarp_vidprom_8b16k_wan_ode_init_1e-5"
# --resume_from_checkpoint "ode_init_diffusers/"
--warp_denoising_step
--log_visualization
--max_train_steps 6001
--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
--dmd_denoising_steps "1000,750,500,250"
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim $NUM_GPUS
--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"
# --init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 500
--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
# --enable_gradient_checkpointing_type "full"
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
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/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,131 @@
#!/bin/bash
#SBATCH --job-name=ode_vidprom2k
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=ode_vidprom2k_output/ode_vidprom2k.out
#SBATCH --error=ode_vidprom2k_output/ode_vidprom2k.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv2
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 WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="/mnt/sharefs/users/hao.zhang/klin/preproc/data/test-ode-preprocessing/"
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/ode_init/validation.json"
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_vidprom2k"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "vidprom2k_wan_ode_init_5e-6"
# --resume_from_checkpoint "ode_init_diffusers/"
--max_train_steps 6001
--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
# --enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 8
--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 100
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-6
--mixed_precision "bf16"
--checkpointing_steps 2000
--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
--enable_gradient_checkpointing_type "full"
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
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/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,132 @@
#!/bin/bash
#SBATCH --job-name=ode_crush
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=ode_crush_output/ode_crush.out
#SBATCH --error=ode_crush_output/ode_crush.err
#SBATCH --exclusive
set -e -x
# Environment Setup
source ~/conda/miniconda/bin/activate
conda activate will-fv2
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 WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
echo "MASTER_ADDR: $MASTER_ADDR"
echo "NODE_RANK: $NODE_RANK"
# MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init_5/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="examples/training/consistency_finetune/ode_init/validation.json"
NUM_GPUS=2
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_warp_2"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "2warp_fixed_wan_ode_init_5e-6"
# --resume_from_checkpoint "ode_init_diffusers/"
# --warp_denoising_step
--max_train_steps 6001
--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
# --enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim $NUM_GPUS
--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 20
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-6
--mixed_precision "bf16"
--checkpointing_steps 2000
--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
--enable_gradient_checkpointing_type "full"
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
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/ode_causal_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,98 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="/mnt/weka/home/hao.zhang/wl/FastVideo2/data/crush-smol_processed_t2v_1_3b_ode_init_single"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
# export CUDA_VISIBLE_DEVICES=4,5
# 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 "overfitwan_ode_init_crush_smol"
# --resume_from_checkpoint "ode_init_diffusers/"
--max_train_steps 2001
# --warp_denoising_step
--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
# --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 20
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--mixed_precision "bf16"
--checkpointing_steps 500
--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
--enable_gradient_checkpointing_type "full"
)
# 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[@]}"
@@ -0,0 +1,24 @@
#!/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_processed_t2v_1_3b_ode_init_single/"
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 "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
}
]
}
@@ -6,8 +6,8 @@ export TOKENIZERS_PARALLELISM=false
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
DATA_DIR="data/crush-smol_processed_t2v_old"
VALIDATION_DATASET_FILE="examples/training/finetune/Wan2.1-Fun-1.3B-InP/crush_smol/validation.json"
NUM_GPUS=4
# export CUDA_VISIBLE_DEVICES=4,5
@@ -52,7 +52,7 @@ dataset_args=(
validation_args=(
--log_validation
--validation_dataset_file $VALIDATION_DATASET_FILE
--validation_steps 200
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
@@ -4,7 +4,7 @@ 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/"
OUTPUT_DIR="data/crush-smol_processed_t2v_old/"
torchrun --nproc_per_node=$GPU_NUM \
fastvideo/pipelines/preprocess/v1_preprocess.py \
@@ -1,6 +1,6 @@
#!/bin/bash
GPU_NUM=2 # 2,4,8
GPU_NUM=1 # 2,4,8
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATASET_PATH="data/crush-smol/"
OUTPUT_DIR="data/crush-smol_processed_t2v/"
@@ -14,7 +14,7 @@ torchrun --nproc_per_node=$GPU_NUM \
--preprocess.dataset_type merged \
--preprocess.dataset_path $DATASET_PATH \
--preprocess.dataset_output_dir $OUTPUT_DIR \
--preprocess.preprocess_video_batch_size 2 \
--preprocess.preprocess_video_batch_size 8 \
--preprocess.dataloader_num_workers 0 \
--preprocess.max_height 480 \
--preprocess.max_width 832 \
@@ -0,0 +1,94 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_old"
VALIDATION_DATASET_FILE="examples/datasets/crush_smol/validation.json"
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# Training arguments
training_args=(
--tracker_project_name "wan_t2v_i2v_finetune"
--output_dir "checkpoints/wan_t2v_i2v_finetune"
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 2
--num_latent_t 20
--num_height 480
--num_width 832
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 4
--tp_size 1
--hsdp_replicate_dim 2
--hsdp_shard_dim 4
)
# 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 5e-5
--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
--enable_gradient_checkpointing_type "full"
--t2v_as_i2v_task True
# --resume_from_checkpoint "checkpoints/wan_t2v_finetune/checkpoint-2500"
)
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wan_t2v_i2v_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,24 @@
#!/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/merge.txt"
OUTPUT_DIR="data/crush-smol_processed_t2v_i2v_1_3b/"
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 2 \
--seed 42 \
--max_height 480 \
--max_width 832 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--train_fps 16 \
--samples_per_file 8 \
--flush_frequency 8 \
--video_length_tolerance_range 5 \
--preprocess_task "t2v_ode_trajectory"
+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
}
+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
@@ -92,6 +92,9 @@ class WanVideoArchConfig(DiTArchConfig):
pos_embed_seq_len: int | None = None
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
# Wan MoE
boundary_ratio: float | None = None
# Causal Wan
local_attn_size: int = -1 # Window size for temporal local attention (-1 indicates global attention)
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
+22 -3
View File
@@ -45,6 +45,8 @@ class PipelineConfig:
embedded_cfg_scale: float = 6.0
flow_shift: float | None = None
disable_autocast: bool = False
ti2v_task: bool = False
t2v_as_i2v_task: bool = False
# Model configuration
dit_config: DiTConfig = field(default_factory=DiTConfig)
@@ -85,9 +87,6 @@ class PipelineConfig:
# DMD parameters
dmd_denoising_steps: list[int] | None = field(default=None)
# Wan2.2 TI2V parameters
ti2v_task: bool = False
# Compilation
# enable_torch_compile: bool = False
@@ -214,6 +213,24 @@ class PipelineConfig:
"Comma-separated list of denoising steps (e.g., '1000,757,522')",
)
# TI2V task
parser.add_argument(
f"--{prefix_with_dot}ti2v-task",
action=StoreBoolean,
dest=f"{prefix_with_dot.replace('-', '_')}ti2v_task",
default=PipelineConfig.ti2v_task,
help="Enable TI2V",
)
# T2V to I2V task
parser.add_argument(
f"--{prefix_with_dot}t2v-as-i2v-task",
action=StoreBoolean,
dest=f"{prefix_with_dot.replace('-', '_')}t2v_as_i2v_task",
default=PipelineConfig.t2v_as_i2v_task,
help="Enable T2V to I2V task",
)
# Add VAE configuration arguments
from fastvideo.configs.models.vaes.base import VAEConfig
VAEConfig.add_cli_args(parser, prefix=f"{prefix_with_dot}vae-config")
@@ -245,7 +262,9 @@ class PipelineConfig:
"""
from fastvideo.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
logger.info("WTF model_path: %s", model_path)
pipeline_config_cls = get_pipeline_config_cls_from_name(model_path)
logger.info("pipeline_config_cls: %s", pipeline_config_cls)
return cast(PipelineConfig, pipeline_config_cls(model_path=model_path))
+4 -2
View File
@@ -13,7 +13,7 @@ from fastvideo.configs.pipelines.wan import (
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
SelfForcingWanT2V480PConfig, Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config,
Wan2_2_TI2V_5B_Config, WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig,
WanT2V720PConfig)
WanT2V720PConfig, SelfForcingWanT2V480PConfig)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
@@ -49,6 +49,7 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
"stepvideo": lambda id: "stepvideo" in id.lower(),
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -60,7 +61,8 @@ 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,
"stepvideo": StepVideoT2VConfig
"stepvideo": StepVideoT2VConfig,
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
# Other fallbacks by architecture
}
+28 -19
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
@@ -82,7 +82,7 @@ class WanI2V480PConfig(WanT2V480PConfig):
default_factory=CLIPVisionConfig)
image_encoder_precision: str = "fp32"
def __post_init__(self):
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@@ -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,40 +104,47 @@ 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])
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
@dataclass
class Wan2_2_TI2V_5B_Config(WanT2V480PConfig):
flow_shift: int = 5
flow_shift: float | None = 5.0
ti2v_task: bool = True
expand_timesteps: bool = True
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
self.dit_config.expand_timesteps = self.expand_timesteps
@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])
@dataclass
class Wan2_2_T2V_A14B_Config(WanT2V480PConfig):
pass
flow_shift: float | None = 12.0
boundary_ratio: float | None = 0.875
def __post_init__(self) -> None:
self.dit_config.boundary_ratio = self.boundary_ratio
@dataclass
class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
pass
class Wan2_2_I2V_A14B_Config(WanI2V480PConfig):
flow_shift: float | None = 5.0
boundary_ratio: float | None = 0.900
def __post_init__(self) -> None:
super().__post_init__()
self.dit_config.boundary_ratio = self.boundary_ratio
# =============================================
@@ -146,5 +153,7 @@ class Wan2_2_I2V_A14B_Config(WanT2V480PConfig):
@dataclass
class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
is_causal: bool = True
flow_shift: float | None = 5.0
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 750, 500, 250])
warp_denoising_step: bool = True
+28
View File
@@ -40,6 +40,7 @@ class SamplingParam:
num_inference_steps: int = 50
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
boundary_ratio: float | None = None
# TeaCache parameters
enable_teacache: bool = False
@@ -47,6 +48,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"
@@ -167,6 +170,12 @@ class SamplingParam:
default=SamplingParam.guidance_rescale,
help="Guidance rescale factor",
)
parser.add_argument(
"--boundary-ratio",
type=float,
default=SamplingParam.boundary_ratio,
help="Boundary timestep ratio",
)
parser.add_argument(
"--save-video",
action="store_true",
@@ -191,6 +200,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
+8 -4
View File
@@ -144,18 +144,22 @@ class Wan2_2_TI2V_5B_SamplingParam(Wan2_2_Base_SamplingParam):
@dataclass
class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
guidance_scale: float = 4.0
guidance_scale_2: float = 3.0
guidance_scale: float = 4.0 # high_noise
guidance_scale_2: float = 3.0 # low_noise
num_inference_steps: int = 40
fps: int = 16
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
# can be overridden during sampling
@dataclass
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
guidance_scale: float = 3.5
guidance_scale_2: float = 3.5
guidance_scale: float = 3.5 # high_noise
guidance_scale_2: float = 3.5 # low_noise
num_inference_steps: int = 40
fps: int = 16
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
# can be overridden during sampling
# =============================================
+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
@@ -39,7 +39,13 @@ def getdataset(args) -> VideoCaptionMergedDataset:
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
@@ -1,5 +1,7 @@
from typing import Any
import numpy as np
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
@@ -120,3 +122,69 @@ def i2v_record_creator(batch: PreprocessBatch) -> list[dict[str, Any]]:
})
return records
def ode_text_only_record_creator(
video_name: str, text_embedding: np.ndarray, caption: str,
trajectory_latents: np.ndarray,
trajectory_timesteps: np.ndarray) -> dict[str, Any]:
"""Create a text-only ODE trajectory record matching pyarrow_schema_ode_trajectory_text_only.
Args:
video_name: Base name/id for the sample (without extension).
text_embedding: Text encoder output array [SeqLen, Dim].
caption: Original text prompt.
trajectory_latents: Collected trajectory latents array.
trajectory_timesteps: Collected timesteps array.
Returns:
dict suitable for records_to_table(…, pyarrow_schema_ode_trajectory_text_only)
"""
assert trajectory_latents is not None, "trajectory_latents is required"
assert trajectory_timesteps is not None, "trajectory_timesteps is required"
record = {
"id": f"text_{video_name}",
"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": caption,
"media_type": "text",
}
record.update({
"trajectory_latents_bytes": trajectory_latents.tobytes(),
"trajectory_latents_shape": list(trajectory_latents.shape),
"trajectory_latents_dtype": str(trajectory_latents.dtype),
})
record.update({
"trajectory_timesteps_bytes": trajectory_timesteps.tobytes(),
"trajectory_timesteps_shape": list(trajectory_timesteps.shape),
"trajectory_timesteps_dtype": str(trajectory_timesteps.dtype),
})
return record
def text_only_record_creator(text_name: str, text_embedding: np.ndarray,
caption: str) -> dict[str, Any]:
"""Create a text-only record matching pyarrow_schema_text_only.
Args:
text_name: Base id/name for the text sample.
text_embedding: Text encoder output array [SeqLen, Dim].
caption: Original text prompt.
Returns:
dict suitable for records_to_table(…, pyarrow_schema_text_only)
"""
record = {
"id": f"text_{text_name}",
"text_embedding_bytes": text_embedding.tobytes(),
"text_embedding_shape": list(text_embedding.shape),
"text_embedding_dtype": str(text_embedding.dtype),
"caption": caption,
}
return record
+38
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,40 @@ 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
])
pyarrow_schema_text_only = pa.schema([
pa.field("id", pa.string()),
# --- Text encoder output tensor ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("text_embedding_bytes", pa.binary()),
# e.g., [SeqLen, Dim]
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
# --- Metadata ---
pa.field("caption", pa.string()),
])
+43
View File
@@ -0,0 +1,43 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.dataset.lmdb_utils import get_array_shape_from_lmdb, retrieve_row_from_lmdb
from torch.utils.data import Dataset
import numpy as np
import torch
import lmdb
# from Self-Forcing: https://github.com/guandeh17/Self-Forcing/blob/main/utils/dataset.py
class ODERegressionLMDBDataset(Dataset):
def __init__(self, data_path: str, max_pair: int = int(1e8)):
print(f"data_path: {data_path}")
self.env = lmdb.open(data_path, readonly=True,
lock=False, readahead=False, meminit=False)
self.latents_shape = get_array_shape_from_lmdb(self.env, 'latents')
self.max_pair = max_pair
def __len__(self):
return min(self.latents_shape[0], self.max_pair)
def __getitem__(self, idx):
"""
Outputs:
- prompts: List of Strings
- latents: Tensor of shape (num_denoising_steps, num_frames, num_channels, height, width). It is ordered from pure noise to clean image.
"""
latents = retrieve_row_from_lmdb(
self.env,
"latents", np.float16, idx, shape=self.latents_shape[1:]
)
if len(latents.shape) == 4:
latents = latents[None, ...]
prompts = retrieve_row_from_lmdb(
self.env,
"prompts", str, idx
)
return {
"prompts": prompts,
"ode_latent": torch.tensor(latents, dtype=torch.float32)
}
+75
View File
@@ -0,0 +1,75 @@
# SPDX-License-Identifier: Apache-2.0
# from Self-Forcing: https://github.com/guandeh17/Self-Forcing/blob/main/utils/lmdb.py
import numpy as np
def get_array_shape_from_lmdb(env, array_name):
with env.begin() as txn:
image_shape = txn.get(f"{array_name}_shape".encode()).decode()
image_shape = tuple(map(int, image_shape.split()))
return image_shape
def store_arrays_to_lmdb(env, arrays_dict, start_index=0):
"""
Store rows of multiple numpy arrays in a single LMDB.
Each row is stored separately with a naming convention.
"""
with env.begin(write=True) as txn:
for array_name, array in arrays_dict.items():
for i, row in enumerate(array):
# Convert row to bytes
if isinstance(row, str):
row_bytes = row.encode()
else:
row_bytes = row.tobytes()
data_key = f'{array_name}_{start_index + i}_data'.encode()
txn.put(data_key, row_bytes)
def process_data_dict(data_dict, seen_prompts):
output_dict = {}
all_videos = []
all_prompts = []
for prompt, video in data_dict.items():
if prompt in seen_prompts:
continue
else:
seen_prompts.add(prompt)
video = video.half().numpy()
all_videos.append(video)
all_prompts.append(prompt)
if len(all_videos) == 0:
return {"latents": np.array([]), "prompts": np.array([])}
all_videos = np.concatenate(all_videos, axis=0)
output_dict['latents'] = all_videos
output_dict['prompts'] = np.array(all_prompts)
return output_dict
def retrieve_row_from_lmdb(lmdb_env, array_name, dtype, row_index, shape=None):
"""
Retrieve a specific row from a specific array in the LMDB.
"""
data_key = f'{array_name}_{row_index}_data'.encode()
with lmdb_env.begin() as txn:
row_bytes = txn.get(data_key)
if dtype == str:
array = row_bytes.decode()
else:
array = np.frombuffer(row_bytes, dtype=dtype)
if shape is not None and len(shape) > 0:
array = array.reshape(shape)
return array
+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"]
+4 -1
View File
@@ -3,9 +3,12 @@ from typing import Any, cast
import numpy as np
import torch
from fastvideo.logger import init_logger
logger = init_logger(__name__)
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,
+128 -1
View File
@@ -1,9 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
# Inspired by SGLang: https://github.com/sgl-project/sglang/blob/main/python/sglang/srt/server_args.py
"""The arguments of FastVideo Inference."""
import argparse
import dataclasses
import json
from contextlib import contextmanager
from dataclasses import field
from enum import Enum
@@ -139,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
@@ -154,6 +158,7 @@ class FastVideoArgs:
"transformer": True,
"vae": True,
})
override_transformer_cls_name: str | None = None
# # DMD parameters
# dmd_denoising_steps: List[int] | None = field(default=None)
@@ -166,6 +171,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
@@ -382,6 +397,12 @@ class FastVideoArgs:
default=FastVideoArgs.enable_stage_verification,
help="Enable input/output verification for pipeline stages",
)
parser.add_argument(
"--override-transformer-cls-name",
type=str,
default=FastVideoArgs.override_transformer_cls_name,
help="Override transformer cls name",
)
# Add pipeline configuration arguments
PipelineConfig.add_cli_args(parser)
@@ -591,6 +612,11 @@ 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
@@ -613,6 +639,7 @@ 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
@@ -644,6 +671,7 @@ 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
@@ -664,16 +692,30 @@ 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
intermediate_latents_visualization: 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
same_step_across_blocks: bool = False # Use same exit timestep for all blocks
last_step_only: bool = False # Only use the last timestep for training
context_noise: int = 0 # Context noise level for cache updates
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
@@ -775,6 +817,20 @@ 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,
@@ -845,6 +901,10 @@ 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")
@@ -949,6 +1009,10 @@ 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")
@@ -985,11 +1049,27 @@ 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,
@@ -1006,6 +1086,11 @@ 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,
@@ -1018,6 +1103,48 @@ 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)")
parser.add_argument(
"--same-step-across-blocks",
action=StoreBoolean,
help="Whether to use the same exit timestep for all blocks")
parser.add_argument(
"--last-step-only",
action=StoreBoolean,
help="Whether to only use the last timestep for training")
parser.add_argument("--context-noise",
type=int,
default=TrainingArgs.context_noise,
help="Context noise level for cache updates")
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 + scale) + shift).flatten(1, 2)
else:
modulated = normalized * (1 + 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 + 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 + scale) + shift
if self.compute_dtype == torch.float32:
output = output.to(x.dtype)
return output
+6 -4
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
@@ -77,9 +77,11 @@ class BaseLayerWithLoRA(nn.Module):
lora_A = self.lora_A.to_local()
if not self.merged and not self.disable_lora:
delta = x @ (
self.slice_lora_b_weights(lora_B.to(x, non_blocking=True))
@ self.slice_lora_a_weights(lora_A.to(x, non_blocking=True)))
lora_A_sliced = self.slice_lora_a_weights(
lora_A.to(x, non_blocking=True))
lora_B_sliced = self.slice_lora_b_weights(
lora_B.to(x, non_blocking=True))
delta = x @ lora_A_sliced.T @ lora_B_sliced.T
if self.lora_alpha != self.lora_rank:
delta = delta * (
self.lora_alpha / self.lora_rank # type: ignore
+80 -46
View File
@@ -147,6 +147,9 @@ class CausalWanSelfAttention(nn.Module):
# Assign new keys/values directly up to current_end
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
# kv_cache["k"] = kv_cache["k"].detach()
# kv_cache["v"] = kv_cache["v"].detach()
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
x = self.attn(
@@ -176,7 +179,7 @@ class CausalWanTransformerBlock(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -209,8 +212,7 @@ class CausalWanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 2. Cross-attention
# Only T2V for now
@@ -223,8 +225,7 @@ class CausalWanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -244,27 +245,39 @@ 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 = self.scale_shift_table + temb
# 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)
assert shift_msa.dtype == torch.float32
6, dim=2)
# *_msa.shape: [batch_size, num_frames, 1, inner_dim]
# assert shift_msa.dtype == torch.float32
# logger.info("temb sum: %s, dtype: %s", temb.float().sum().item(), temb.dtype)
# logger.info("scale_msa sum: %s, dtype: %s", scale_msa.float().sum().item(), scale_msa.dtype)
# logger.info("shift_msa sum: %s, dtype: %s", shift_msa.float().sum().item(), shift_msa.dtype)
# 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).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).flatten(1, 2)
# logger.info("norm_hidden_states sum: %s, shape: %s", norm_hidden_states.float().sum().item(), norm_hidden_states.shape)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k(key)
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -278,8 +291,6 @@ class CausalWanTransformerBlock(nn.Module):
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
@@ -288,13 +299,10 @@ class CausalWanTransformerBlock(nn.Module):
crossattn_cache=crossattn_cache)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
@@ -357,8 +365,7 @@ class CausalWanTransformer3DModel(BaseDiT):
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
@@ -368,7 +375,7 @@ class CausalWanTransformer3DModel(BaseDiT):
# Causal-specific
self.block_mask = None
self.num_frame_per_block = 1
self.num_frame_per_block = 3
self.independent_first_frame = False
self.__post_init__()
@@ -480,15 +487,19 @@ class CausalWanTransformer3DModel(BaseDiT):
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
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,19 +537,15 @@ 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)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
output = self.unpatchify(hidden_states, grid_sizes)
return output
return torch.stack(output)
def _forward_train(self,
hidden_states: torch.Tensor,
@@ -579,8 +586,8 @@ class CausalWanTransformer3DModel(BaseDiT):
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
# Construct blockwise causal attn mask
if self.block_mask is None:
@@ -593,11 +600,15 @@ class CausalWanTransformer3DModel(BaseDiT):
)
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
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,19 +634,15 @@ 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)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
output = self.unpatchify(hidden_states, grid_sizes)
return output
return torch.stack(output)
def forward(
self,
@@ -646,3 +653,30 @@ class CausalWanTransformer3DModel(BaseDiT):
return self._forward_inference(*args, **kwargs)
else:
return self._forward_train(*args, **kwargs)
def unpatchify(self, x, grid_sizes):
r"""
Args:
x (List[Tensor]):
List of patchified features, each with shape [L, C_out * prod(patch_size)]
grid_sizes (Tensor):
Original spatial-temporal grid dimensions before patching,
Returns:
Tensor:
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
"""
c = self.out_channels
out = []
for u, v in zip(x, grid_sizes.tolist()):
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
u = u.permute(6, 0, 3, 1, 4, 2, 5)
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
out.append(u)
return out
+80 -73
View File
@@ -1,3 +1,5 @@
import torch
import torch.nn as nn
# SPDX-License-Identifier: Apache-2.0
import math
@@ -37,16 +39,14 @@ class WanImageEmbedding(torch.nn.Module):
def __init__(self, in_features: int, out_features: int):
super().__init__()
self.norm1 = FP32LayerNorm(in_features)
self.norm1 = nn.LayerNorm(in_features)
self.ff = MLP(in_features, in_features, out_features, act_type="gelu")
self.norm2 = FP32LayerNorm(out_features)
self.norm2 = nn.LayerNorm(out_features)
def forward(self,
encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
dtype = encoder_hidden_states_image.dtype
def forward(self, encoder_hidden_states_image: torch.Tensor) -> torch.Tensor:
hidden_states = self.norm1(encoder_hidden_states_image)
hidden_states = self.ff(hidden_states)
hidden_states = self.norm2(hidden_states).to(dtype)
hidden_states = self.norm2(hidden_states)
return hidden_states
@@ -62,7 +62,7 @@ class WanTimeTextImageEmbedding(nn.Module):
super().__init__()
self.time_embedder = TimestepEmbedder(
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
dim, frequency_embedding_size=time_freq_dim, act_layer="silu", freq_dtype=torch.float64)
self.time_modulation = ModulateProjection(dim,
factor=6,
act_layer="silu")
@@ -156,12 +156,12 @@ class WanT2VCrossAttention(WanSelfAttention):
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
if crossattn_cache is not None:
if not crossattn_cache["is_init"]:
crossattn_cache["is_init"] = True
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
crossattn_cache["k"] = k
crossattn_cache["v"] = v
@@ -169,7 +169,7 @@ class WanT2VCrossAttention(WanSelfAttention):
k = crossattn_cache["k"]
v = crossattn_cache["v"]
else:
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
# compute attention
@@ -213,10 +213,10 @@ class WanI2VCrossAttention(WanSelfAttention):
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
q = self.norm_q.forward_native(self.to_q(x)[0]).view(b, -1, n, d)
k = self.norm_k.forward_native(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
k_img = self.norm_added_k.forward_native(self.add_k_proj(context_img)[0]).view(
b, -1, n, d)
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
img_x = self.attn(q, k_img, v_img)
@@ -247,7 +247,7 @@ class WanTransformerBlock(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -278,29 +278,29 @@ class WanTransformerBlock(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
# I2V
self.attn2 = WanI2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
self.attn2 = WanI2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
else:
# T2V
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dim,
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -319,12 +319,11 @@ class WanTransformerBlock(nn.Module):
hidden_states = hidden_states.squeeze(1)
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
if temb.dim() == 4:
# temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
self.scale_shift_table.unsqueeze(0) + temb.float()
self.scale_shift_table.unsqueeze(0) + temb
).chunk(6, dim=2)
# batch_size, seq_len, 1, inner_dim
shift_msa = shift_msa.squeeze(2)
@@ -335,22 +334,20 @@ class WanTransformerBlock(nn.Module):
c_gate_msa = c_gate_msa.squeeze(2)
else:
# temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B)
e = self.scale_shift_table + temb.float()
e = self.scale_shift_table + temb
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
norm_hidden_states = self.norm1(hidden_states) * (1 + scale_msa) + shift_msa
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k(key)
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -370,26 +367,20 @@ class WanTransformerBlock(nn.Module):
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
context=encoder_hidden_states,
attn_output = self.attn2(norm_hidden_states,
context=encoder_hidden_states,
context_lens=None)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
class WanTransformerBlock_VSA(nn.Module):
def __init__(self,
@@ -406,7 +397,7 @@ class WanTransformerBlock_VSA(nn.Module):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.norm1 = nn.LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
@@ -438,8 +429,7 @@ class WanTransformerBlock_VSA(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
@@ -459,8 +449,7 @@ class WanTransformerBlock_VSA(nn.Module):
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
@@ -480,23 +469,22 @@ class WanTransformerBlock_VSA(nn.Module):
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
e = self.scale_shift_table + temb
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
norm_hidden_states = (self.norm1(hidden_states) *
(1 + scale_msa) + shift_msa)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
gate_compress, _ = self.to_gate_compress(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k(key)
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
@@ -521,8 +509,6 @@ class WanTransformerBlock_VSA(nn.Module):
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
@@ -530,17 +516,15 @@ class WanTransformerBlock_VSA(nn.Module):
context_lens=None)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
class WanTransformer3DModel(CachableDiT):
_fsdp_shard_conditions = WanVideoConfig()._fsdp_shard_conditions
_compile_conditions = WanVideoConfig()._compile_conditions
@@ -598,8 +582,7 @@ class WanTransformer3DModel(CachableDiT):
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
@@ -659,10 +642,12 @@ class WanTransformer3DModel(CachableDiT):
rope_theta=10000)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
freqs_sin.float()) if freqs_cos is not None else None
freqs_cis = (freqs_cos,
freqs_sin) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
grid_sizes = torch.stack(
[torch.tensor(hidden_states[0].shape[1:], dtype=torch.long)])
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
@@ -672,6 +657,8 @@ class WanTransformer3DModel(CachableDiT):
else:
ts_seq_len = None
encoder_hidden_states = torch.cat([encoder_hidden_states, encoder_hidden_states.new_zeros(1, self.text_len - encoder_hidden_states.size(1), encoder_hidden_states.size(2))], dim=1)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len)
if ts_seq_len is not None:
@@ -728,14 +715,35 @@ class WanTransformer3DModel(CachableDiT):
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
output = self.unpatchify(hidden_states, grid_sizes)
return output
return torch.stack(output)
def unpatchify(self, x, grid_sizes):
r"""
Args:
x (List[Tensor]):
List of patchified features, each with shape [L, C_out * prod(patch_size)]
grid_sizes (Tensor):
Original spatial-temporal grid dimensions before patching,
Returns:
Tensor:
Reconstructed video tensors with shape [B, C_out, F, H / 8, W / 8]
"""
c = self.out_channels
out = []
for u, v in zip(x, grid_sizes.tolist()):
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
u = u.permute(6, 0, 3, 1, 4, 2, 5)
# u = torch.einsum('fhwpqrc->cfphqwr', u.contiguous())
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size)])
out.append(u)
return out
def maybe_cache_states(self, hidden_states: torch.Tensor,
original_hidden_states: torch.Tensor) -> None:
@@ -827,5 +835,4 @@ class WanTransformer3DModel(CachableDiT):
if self.is_even:
return hidden_states + self.previous_residual_even
else:
return hidden_states + self.previous_residual_odd
return hidden_states + self.previous_residual_odd
@@ -415,6 +415,10 @@ class TransformerLoader(ComponentLoader):
raise ValueError(
"Model config does not contain a _class_name attribute. "
"Only diffusers format is supported.")
logger.info("transformer cls_name: %s", cls_name)
if fastvideo_args.override_transformer_cls_name is not None:
cls_name = fastvideo_args.override_transformer_cls_name
logger.info("Overriding transformer cls_name to %s", cls_name)
fastvideo_args.model_paths["transformer"] = model_path
@@ -430,6 +434,16 @@ 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)
+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 = {
@@ -635,8 +635,31 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
noise: torch.Tensor,
timestep: torch.IntTensor,
) -> torch.Tensor:
"""
Args:
clean_latent: the clean latent with shape [B, C, H, W],
where B is batch_size or batch_size * num_frames
noise: the noise with shape [B, C, H, W]
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
Returns:
the corrupted latent with shape [B, C, H, W]
"""
# If timestep is [bs, num_frames]
if timestep.ndim == 2:
timestep = timestep.flatten(0, 1)
assert timestep.numel() == clean_latent.shape[0]
elif timestep.ndim == 1:
# If timestep is [1]
if timestep.shape[0] == 1:
timestep = timestep.expand(clean_latent.shape[0])
else:
assert timestep.numel() == clean_latent.shape[0]
else:
raise ValueError(f"[add_noise] Invalid timestep shape: {timestep.shape}")
# timestep shape should be [B]
self.sigmas = self.sigmas.to(noise.device)
timestep = timestep.expand(clean_latent.shape[0])
self.timesteps = self.timesteps.to(noise.device)
timestep_id = torch.argmin(
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
@@ -0,0 +1,124 @@
# 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):
config_name = "scheduler_config.json"
order = 1
@register_to_config
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
+30 -5
View File
@@ -145,14 +145,39 @@ 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)
noise_input_latent = noise_input_latent.float().to(device)
sigmas = scheduler.sigmas.float().to(device)
timesteps = scheduler.timesteps.float().to(device)
# Convert to double following Self-Forcing
# https://github.com/guandeh17/Self-Forcing/blob/main/utils/wan_wrapper.py#L184
pred_noise = pred_noise.double().to(device)
noise_input_latent = noise_input_latent.double().to(device)
sigmas = scheduler.sigmas.double().to(device)
timesteps = scheduler.timesteps.double().to(device)
timestep_id = torch.argmin(
(timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
sigma_t = sigmas[timestep_id].reshape(-1, 1, 1, 1)
@@ -16,8 +16,7 @@ from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
CausalDMDDenosingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
TextEncodingStage)
# isort: on
logger = init_logger(__name__)
@@ -29,10 +28,6 @@ class WanCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
"""Set up pipeline stages with proper dependency injection."""
@@ -48,10 +43,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"),
+32 -12
View File
@@ -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)
@@ -121,14 +121,25 @@ class ComposedPipelineBase(ABC):
model_path: str,
device: str | None = None,
torch_dtype: torch.dtype | None = None,
pipeline_config: str | PipelineConfig | None = None,
pipeline_config: PipelineConfig | None = None,
args: argparse.Namespace | None = None,
required_config_modules: list[str] | None = None,
loaded_modules: dict[str, torch.nn.Module]
| None = None,
**kwargs) -> "ComposedPipelineBase":
"""
Load a pipeline from a pretrained model.
Load a pipeline from a pretrained model.
Few different patterns are supported:
- Only provide model_path:
- This will load the pipeline in inference mode.
- The pipeline will be initialized with the default config.
- The pipeline will be initialized with the default modules.
- The pipeline will be initialized with the default stages.
- The pipeline will be initialized with the default stages.
- override the default config using pipeline_config or args or kwargs
- override the default modules using loaded_modules
- override the pipelineconfig
loaded_modules: Optional[Dict[str, torch.nn.Module]] = None,
If provided, loaded_modules will be used instead of loading from config/pretrained weights.
"""
@@ -136,9 +147,18 @@ class ComposedPipelineBase(ABC):
kwargs['model_path'] = model_path
fastvideo_args = FastVideoArgs.from_kwargs(**kwargs)
if pipeline_config is not None:
fastvideo_args.pipeline_config = pipeline_config
if fastvideo_args.override_transformer_cls_name is not None:
pipeline_config = PipelineConfig.from_pretrained("wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers")
fastvideo_args.pipeline_config = pipeline_config
else:
assert args is not None, "args must be provided for training mode"
fastvideo_args = TrainingArgs.from_cli_args(args)
if fastvideo_args.override_transformer_cls_name is not None:
pipeline_config = PipelineConfig.from_pretrained("wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers")
fastvideo_args.pipeline_config = pipeline_config
logger.info("in 2 Overriding transformer cls name to %s", fastvideo_args.override_transformer_cls_name)
# TODO(will): fix this so that its not so ugly
fastvideo_args.model_path = model_path
for key, value in kwargs.items():
@@ -149,7 +169,8 @@ class ComposedPipelineBase(ABC):
# model is loaded with the correct precision. Subsequently we will
# use FSDP2's MixedPrecisionPolicy to set the precision for the
# fwd, bwd, and other operations' precision.
assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
fastvideo_args.pipeline_config.dit_precision = 'fp32'
# assert fastvideo_args.pipeline_config.dit_precision == 'fp32', 'only fp32 is supported for training'
logger.info("fastvideo_args in from_pretrained: %s", fastvideo_args)
@@ -237,20 +258,19 @@ class ComposedPipelineBase(ABC):
# remove keys that are not pipeline modules
model_index.pop("_class_name")
model_index.pop("_diffusers_version")
# @TODO(Wei): Temporary hack
if "boundary_ratio" in model_index and model_index[
"boundary_ratio"] is not None:
logger.info(
"MoE pipeline detected. Adding transformer_2 to self.required_config_modules..."
)
self.required_config_modules.append("transformer_2")
if fastvideo_args.boundary_ratio is None:
logger.info(
"MoE pipeline detected. Setting boundary ratio to %s",
model_index["boundary_ratio"])
fastvideo_args.boundary_ratio = model_index["boundary_ratio"]
logger.info("MoE pipeline detected. Setting boundary ratio to %s",
model_index["boundary_ratio"])
fastvideo_args.pipeline_config.dit_config.boundary_ratio = model_index[
"boundary_ratio"]
model_index.pop("boundary_ratio", None)
# used by Wan2.2 ti2v
model_index.pop("expand_timesteps", None)
# some sanity checks
@@ -283,8 +303,8 @@ class ComposedPipelineBase(ABC):
architecture) in model_index.items():
if transformers_or_diffusers is None:
logger.warning(
"Module in model_index.json has null value, removing from required_config_modules"
)
"Module %s in model_index.json has null value, removing from required_config_modules",
module_name)
if module_name in self.required_config_modules:
self.required_config_modules.remove(module_name)
continue
+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 -1
View File
@@ -129,6 +129,7 @@ class ForwardBatch:
timesteps: torch.Tensor | None = None
timestep: torch.Tensor | float | int | None = None
step_index: int | None = None
boundary_ratio: float | None = None
# Scheduler parameters
num_inference_steps: int = 50
@@ -147,7 +148,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[int] | 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 +212,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
@@ -236,6 +246,7 @@ class TrainingBatch:
fake_score_loss: float = 0.0
dmd_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
latent_vis_dict: dict[str, torch.Tensor] = field(default_factory=dict)
fake_score_latent_vis_dict: dict[str, Any] = field(default_factory=dict)
@@ -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
@@ -12,6 +10,8 @@ from torch.utils.data import DataLoader
from tqdm import tqdm
from fastvideo.dataset import getdataset
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
records_to_table)
from fastvideo.dataset.preprocessing_datasets import PreprocessBatch
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
@@ -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,22 @@ 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
@@ -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,654 @@
# 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 PIL import Image
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm import tqdm
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import getdataset
from fastvideo.dataset.dataloader.schema import pyarrow_schema_ode_trajectory
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
ImageVAEEncodingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
from fastvideo.utils import save_decoded_latents_as_video, shallow_asdict
logger = init_logger(__name__)
class FlowMatchScheduler:
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):
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)
def set_timesteps(self,
num_inference_steps=100,
denoising_strength=1.0,
training=False,
device=None):
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,
timestep,
sample,
to_final=False,
return_dict=False,
**kwargs):
assert return_dict is False
assert kwargs == {}
self.sigmas = self.sigmas.to(model_output.device)
self.timesteps = self.timesteps.to(model_output.device)
logger.info('step timestep: %s', timestep)
logger.info('step timestep: %s', timestep.shape)
# timestep is [num_frames]
# timestep_id = torch.argmin(
# (self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
# assert timestep.ndim == 1
# assert timestep.shape[0] == 1
timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0)
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)
return (prev_sample, )
def scale_model_input(self, sample: torch.Tensor, *args,
**kwargs) -> torch.Tensor:
"""
Ensures interchangeability with schedulers that need to scale the denoising model input depending on the
current timestep.
Args:
sample (`torch.Tensor`):
The input sample.
Returns:
`torch.Tensor`:
A scaled input sample.
"""
return sample
def add_noise(self, original_samples, noise, timestep):
"""
Diffusion forward corruption process.
Input:
- clean_latent: the clean latent with shape [B, C, H, W]
- noise: the noise with shape [B, C, H, W]
- timestep: the timestep with shape [B]
Output: the corrupted latent with shape [B, C, H, W]
"""
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):
timestep_id = torch.argmin(
(self.timesteps - timestep.to(self.timesteps.device)).abs())
weights = self.linear_timesteps_weights[timestep_id]
return weights
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]]
def get_schema_fields(self):
"""Get the schema fields for ODE Trajectory pipeline."""
return [f.name for f in pyarrow_schema_ode_trajectory]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
fastvideo_args.pipeline_config.flow_shift = 5
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"] = FlowMatchScheduler(
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)
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="vae_encoding_stage",
stage=ImageVAEEncodingStage(
vae=self.get_module("vae"), ))
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
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"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
pipeline=self,
))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
def preprocess_video_and_text_and_trajectory(self,
fastvideo_args: FastVideoArgs,
args):
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
# Filter out invalid samples (those with all zeros)
valid_indices = []
for i, pixel_values in enumerate(data["pixel_values"]):
if not torch.all(
pixel_values == 0): # Check if all values are zero
valid_indices.append(i)
self.num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples
valid_data = {
"pixel_values":
torch.stack(
[data["pixel_values"][i] for i in valid_indices]),
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
"fps": [data["fps"][i] for i in valid_indices],
"duration": [data["duration"][i] for i in valid_indices],
}
# VAE
with torch.autocast("cuda", dtype=torch.float32):
latents = self.get_module("vae").encode(
valid_data["pixel_values"].to(
get_local_torch_device())).mean
# Get extra features if needed
extra_features = self.get_extra_features(
valid_data, fastvideo_args)
batch_captions = valid_data["text"]
logger.info(f"===== batch_captions: {batch_captions}")
# Encode text using the standalone TextEncodingStage API
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
# # Get sequence lengths from attention masks (number of 1s)
# seq_lens = prompt_attention_mask.sum(dim=1)
# non_padded_embeds = []
# non_padded_masks = []
# # Process each item in the batch
# for i in range(prompt_embeds.size(0)):
# seq_len = seq_lens[i].item()
# # Slice the embeddings and masks to keep only non-padding parts
# non_padded_embeds.append(prompt_embeds[i, :seq_len])
# non_padded_masks.append(prompt_attention_mask[i, :seq_len])
# Update the tensors with non-padded versions
# prompt_embeds = non_padded_embeds
# prompt_attention_masks = non_padded_masks
# prompt_embeds = prompt_embeds
# logger.info(f"===== prompt_embeds: {prompt_embeds[0].shape}")
# logger.info(f"===== prompt_attention_masks: {prompt_attention_masks[0].shape}")
sampling_params = SamplingParam.from_pretrained(args.model_path)
# encode negative prompt for trajectory collection
if sampling_params.guidance_scale > 1 and sampling_params.negative_prompt is not None:
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
sampling_params.negative_prompt,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
negative_prompt_embed = negative_prompt_embeds_list[0][0]
negative_prompt_attention_mask = negative_prompt_masks_list[
0][0]
else:
negative_prompt_embed = None
negative_prompt_attention_mask = None
trajectory_latents = []
trajectory_timesteps = []
trajectory_decoded = []
for i, (prompt_embed, prompt_attention_mask) in enumerate(
zip(prompt_embeds, prompt_attention_masks, strict=False)):
prompt_embed = prompt_embed.unsqueeze(0)
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
logger.info("what")
logger.info(f"===== prompt_embed: {prompt_embed.shape}")
logger.info(
f"===== prompt_attention_mask: {prompt_attention_mask.shape}"
)
# Collect the trajectory data
batch = ForwardBatch(
**shallow_asdict(sampling_params),
# data_type="video",
# seed=args.seed,
# prompt=batch_captions[i],
# prompt_embeds=[prompt_embed],
# prompt_attention_mask=[prompt_attention_mask],
# height=args.max_height,
# width=args.max_width,
# num_frames=81,
# fps=args.train_fps,
# return_trajectory_latents=True,
# guidance_scale=3.0,
# do_classifier_free_guidance=True,
)
batch.prompt_embeds = [prompt_embed]
batch.prompt_attention_mask = [prompt_attention_mask]
batch.negative_prompt_embeds = [negative_prompt_embed]
batch.negative_attention_mask = [
negative_prompt_attention_mask
]
batch.return_trajectory_latents = True
batch.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
batch.num_inference_steps = 48
# batch.num_frames = 81
batch.fps = args.train_fps
batch.guidance_scale = 6.0
batch.do_classifier_free_guidance = True
# fastvideo_args.pipeline_config.ti2v_task = True
result_batch = self.input_validation_stage(
batch, fastvideo_args)
# result_batch = self.prompt_encoding_stage(result_batch, fastvideo_args)
# result_batch = self.vae_encoding_stage(result_batch, fastvideo_args)
result_batch = self.timestep_preparation_stage(
batch, fastvideo_args)
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch,
fastvideo_args)
# trajectory_latents = result_batch.trajectory_latents
trajectory_latents.append(
result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(
result_batch.trajectory_timesteps.cpu())
trajectory_decoded.append(result_batch.trajectory_decoded)
extra_features["trajectory_latents"] = trajectory_latents
extra_features["trajectory_timesteps"] = trajectory_timesteps
logger.info(
f"===== trajectory_latents: {trajectory_latents[0].shape}")
logger.info(
f"===== trajectory_latents len: {len(trajectory_latents)}")
logger.info(f"===== trajectory_timesteps: {trajectory_timesteps}")
logger.info(
f"===== trajectory_timesteps len: {len(trajectory_timesteps)}")
if batch.return_trajectory_decoded:
logger.info("===== SAVING TRAJECTORY DECODED")
for i, decoded_frames in enumerate(trajectory_decoded):
for j, decoded_frame in enumerate(decoded_frames):
logger.info(
f"===== SAVING TRAJECTORY DECODED {i} for prompt {batch_captions[i]}"
)
save_decoded_latents_as_video(
decoded_frame,
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
args.train_fps)
# assert False
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, video_path in save_pbar:
# Get the corresponding latent and info using video name
latent = latents[idx].cpu()
video_name = os.path.basename(video_path).split(".")[0]
# Convert tensors to numpy arrays
vae_latent = latent.cpu().numpy()
text_embedding = prompt_embeds[idx].cpu().numpy()
# Get extra features for this sample if needed
sample_extra_features = {}
if extra_features:
for key, value in extra_features.items():
logger.info(f"===== key: {key}")
if isinstance(value, torch.Tensor):
logger.info(f"===== value: {value[idx].shape}")
sample_extra_features[key] = value[idx].cpu().numpy(
)
else:
assert isinstance(value, list)
if isinstance(value[idx], torch.Tensor):
logger.info(
f"===== value in list: {value[idx].shape}")
sample_extra_features[key] = value[idx].cpu(
).float().numpy()
else:
logger.info("===== value in list: not tensor")
sample_extra_features[key] = value[idx]
# logger.info(f"===== value: not tensor")
# sample_extra_features[key] = value[idx]
# Create record for Parquet dataset
record = self.create_record(
video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=sample_extra_features)
batch_data.append(record)
if batch_data:
# Add progress bar for writing to Parquet dataset
write_pbar = tqdm(total=1,
desc="Writing to Parquet dataset",
unit="batch")
# Convert batch data to PyArrow arrays
arrays = []
for field in self.get_schema_fields():
if field.endswith('_bytes'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.binary()))
elif field.endswith('_shape'):
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.list_(pa.int32())))
elif field in ['width', 'height', 'num_frames']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.int32()))
elif field in ['duration_sec', 'fps']:
arrays.append(
pa.array([record[field] for record in batch_data],
type=pa.float32()))
else:
arrays.append(
pa.array([record[field] for record in batch_data]))
table = pa.Table.from_arrays(arrays,
names=self.get_schema_fields())
write_pbar.update(1)
write_pbar.close()
# Store the table in a list for later processing
if not hasattr(self, 'all_tables'):
self.all_tables = []
self.all_tables.append(table)
logger.info("Collected batch with %s samples", len(table))
if self.num_processed_samples >= args.flush_frequency:
self._flush_tables(self.num_processed_samples, args,
self.combined_parquet_dir)
self.num_processed_samples = 0
self.all_tables = []
def get_extra_features(self, valid_data: dict[str, Any],
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
# TODO(will): move these to cpu at some point
self.get_module("vae").to(get_local_torch_device())
# generator = torch.Generator(device=get_local_torch_device(), seed=42)
generator = torch.Generator("cpu").manual_seed(42)
features = {}
"""Get CLIP features from the first frame of each video."""
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
_, _, num_frames, height, width = valid_data["pixel_values"].shape
# latent_height = height // self.get_module(
# "vae").spatial_compression_ratio
# latent_width = width // self.get_module("vae").spatial_compression_ratio
unprocessed_images = []
pil_images = []
# Frame has values between -1 and 1
for frame in first_frame:
frame = (frame + 1) * 127.5
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
pil_images.append(frame_pil)
# processed_img = self.get_module("image_processor")(
# images=frame_pil, return_tensors="pt")
unprocessed_images.append(frame_pil)
"""Get VAE features from the first frame of each video"""
video_conditions = []
for frame in unprocessed_images:
latent = self.vae_encoding_stage.encode_image(
frame, height, width, fastvideo_args, generator)
video_conditions.append(latent)
features["image_condition_latents"] = video_conditions
features["pil_images"] = pil_images
return features
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
valid_data: dict[str, Any],
idx: int,
extra_features: dict[str, Any] | None = None) -> dict[str, Any]:
"""Create a record for the Parquet dataset with CLIP features."""
record = super().create_record(video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
if extra_features and "image_condition_latents" in extra_features:
image_condition_latents = extra_features["image_condition_latents"]
record.update({
"image_condition_latents_bytes":
image_condition_latents.tobytes(),
"image_condition_latents_shape":
list(image_condition_latents.shape),
"image_condition_latents_dtype":
str(image_condition_latents.dtype),
})
else:
record.update({
"image_condition_latents_bytes": b"",
"image_condition_latents_shape": [],
"image_condition_latents_dtype": "",
})
if extra_features and "trajectory_latents" in extra_features:
trajectory_latents = extra_features["trajectory_latents"]
record.update({
"trajectory_latents_bytes":
trajectory_latents.tobytes(),
"trajectory_latents_shape":
list(trajectory_latents.shape),
"trajectory_latents_dtype":
str(trajectory_latents.dtype),
})
else:
record.update({
"trajectory_latents_bytes": b"",
"trajectory_latents_shape": [],
"trajectory_latents_dtype": "",
})
if extra_features and "trajectory_timesteps" in extra_features:
trajectory_timesteps = extra_features["trajectory_timesteps"]
record.update({
"trajectory_timesteps_bytes":
trajectory_timesteps.tobytes(),
"trajectory_timesteps_shape":
list(trajectory_timesteps.shape),
"trajectory_timesteps_dtype":
str(trajectory_timesteps.dtype),
})
else:
record.update({
"trajectory_timesteps_bytes": b"",
"trajectory_timesteps_shape": [],
"trajectory_timesteps_dtype": "",
})
if extra_features and "pil_image" in extra_features:
pil_image = extra_features["pil_image"]
record.update({
"pil_image_bytes": pil_image.tobytes(),
"pil_image_shape": list(pil_image.shape),
"pil_image_dtype": str(pil_image.dtype),
})
else:
record.update({
"pil_image_bytes": b"",
"pil_image_shape": [],
"pil_image_dtype": "",
})
return record
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
if not self.post_init_called:
self.post_init()
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 = getdataset(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_video_and_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
@@ -0,0 +1,184 @@
# SPDX-License-Identifier: Apache-2.0
"""
Text-only Data Preprocessing pipeline implementation.
This module contains an implementation of the Text-only Data Preprocessing pipeline
using the modular pipeline architecture, based on the ODE Trajectory preprocessing.
"""
import os
from collections.abc import Iterator
from typing import Any
import torch
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm import tqdm
from fastvideo.dataset import gettextdataset
from fastvideo.dataset.dataloader.parquet_io import (ParquetDatasetWriter,
records_to_table)
from fastvideo.dataset.dataloader.record_schema import text_only_record_creator
from fastvideo.dataset.dataloader.schema import pyarrow_schema_text_only
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import TextEncodingStage
logger = init_logger(__name__)
class PreprocessPipeline_Text(BasePreprocessPipeline):
"""Text-only preprocessing pipeline implementation."""
_required_config_modules = ["text_encoder", "tokenizer"]
preprocess_dataloader: StatefulDataLoader
preprocess_loader_iter: Iterator[dict[str, Any]]
pbar: Any
num_processed_samples: int = 0
def get_pyarrow_schema(self):
"""Return the PyArrow schema for text-only pipeline."""
return pyarrow_schema_text_only
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
def preprocess_text_only(self, fastvideo_args: FastVideoArgs, args):
"""Preprocess text-only data."""
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
for i, text in enumerate(data["text"]):
if text and text.strip(): # Check if text is not empty
valid_indices.append(i)
self.num_processed_samples += len(valid_indices)
if not valid_indices:
continue
# Create new batch with only valid samples (text-only)
valid_data = {
"text": [data["text"][i] for i in valid_indices],
"path": [data["path"][i] for i in valid_indices],
}
batch_captions = valid_data["text"]
# Encode text using the standalone TextEncodingStage API
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=[0],
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
logger.info("===== prompt_embeds: %s", prompt_embeds.shape)
logger.info("===== prompt_attention_masks: %s",
prompt_attention_masks.shape)
# Prepare batch data for Parquet dataset
batch_data = []
# Add progress bar for saving outputs
save_pbar = tqdm(enumerate(valid_data["path"]),
desc="Saving outputs",
unit="item",
leave=False)
for idx, text_path in save_pbar:
text_name = os.path.basename(text_path).split(".")[0]
# Convert tensors to numpy arrays
text_embedding = prompt_embeds[idx].cpu().numpy()
# Create record for Parquet dataset (text-only schema)
record = text_only_record_creator(
text_name=text_name,
text_embedding=text_embedding,
caption=valid_data["text"][idx],
)
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,
pyarrow_schema_text_only)
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(write_remainder=True)
if written:
logger.info("Final flush wrote %s samples", written)
# Text-only record creation moved to fastvideo.dataset.dataloader.record_schema
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs, args):
if not self.post_init_called:
self.post_init()
self.local_rank = int(os.getenv("RANK", 0))
os.makedirs(args.output_dir, exist_ok=True)
# Create directory for combined data
self.combined_parquet_dir = os.path.join(args.output_dir,
"combined_parquet_dataset")
os.makedirs(self.combined_parquet_dir, exist_ok=True)
# Loading text dataset
train_dataset = gettextdataset(args)
self.preprocess_dataloader = DataLoader(
train_dataset,
batch_size=args.preprocess_video_batch_size,
num_workers=args.dataloader_num_workers,
)
self.preprocess_loader_iter = iter(self.preprocess_dataloader)
self.num_processed_samples = 0
# Add progress bar for text preprocessing
self.pbar = tqdm(self.preprocess_loader_iter,
desc="Processing text",
unit="batch",
disable=self.local_rank != 0)
# Initialize class variables for data sharing
self.text_data: dict[str, Any] = {} # Store text metadata and paths
self.preprocess_text_only(fastvideo_args, args)
EntryClass = PreprocessPipeline_Text
@@ -1,5 +1,6 @@
import argparse
import os
from typing import Any
from fastvideo import PipelineConfig
from fastvideo.configs.models.vaes import WanVAEConfig
@@ -9,8 +10,12 @@ 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.pipelines.preprocess.preprocess_pipeline_text import (
PreprocessPipeline_Text)
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
@@ -21,12 +26,22 @@ def main(args) -> None:
maybe_init_distributed_environment_and_model_parallel(1, 1)
num_gpus = int(os.environ["WORLD_SIZE"])
assert num_gpus == 1, "Only support 1 GPU"
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
kwargs = {
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=False),
}
kwargs: dict[str, Any] = {}
if args.preprocess_task == "text_only":
kwargs = {
"text_encoder_cpu_offload": False,
}
else:
# Full config for video/image processing
kwargs = {
"vae_precision": "fp32",
"vae_config": WanVAEConfig(load_encoder=True, load_decoder=True),
}
pipeline_config.update_config_from_dict(kwargs)
fastvideo_args = FastVideoArgs(
model_path=args.model_path,
num_gpus=get_world_size(),
@@ -35,7 +50,19 @@ 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 == "text_only":
PreprocessPipeline = PreprocessPipeline_Text
else:
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}. "
f"Valid options: t2v, i2v, ode_trajectory, text_only")
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)
@@ -74,7 +101,11 @@ if __name__ == "__main__":
parser.add_argument("--video_length_tolerance_range", type=int, default=2.0)
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--preprocess_task", type=str, default="t2v")
parser.add_argument("--preprocess_task",
type=str,
default="t2v",
choices=["t2v", "i2v", "text_only"],
help="Type of preprocessing task to run")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
parser.add_argument("--text_max_length", type=int, default=256)
+45 -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,18 @@ 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]
else:
assert False, "warp_denoising_step must be true"
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 +271,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 +339,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 +421,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
+63 -45
View File
@@ -50,6 +50,50 @@ 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 latents into pixel space."""
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,
@@ -66,6 +110,7 @@ class DecodingStage(PipelineStage):
Returns:
The batch with decoded outputs.
"""
# load vae if not already loaded (used for memory constrained devices)
pipeline = self.pipeline() if self.pipeline else None
if not fastvideo_args.model_loaded["vae"]:
loader = VAELoader()
@@ -75,58 +120,31 @@ 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 = []
logger.info(f"batch.trajectory_latents.shape: {batch.trajectory_latents.shape}")
assert batch.trajectory_latents is not None, "batch should have trajectory latents"
for idx in range(batch.trajectory_latents.shape[1]):
# bathc.trajectory_latents is [batch_size, timesteps, channels, frames, height, width]
cur_latent = batch.trajectory_latents[:, idx, :, :, :, :]
logger.info(f"cur_latent.shape: {cur_latent.shape}")
cur_timestep = batch.trajectory_timesteps[idx]
logger.info(
f"decoding trajectory latent for timestep: {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'):
+114 -14
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
)
@@ -132,11 +140,12 @@ class DenoisingStage(PipelineStage):
latents = latents[:, :, rank_in_sp_group, :, :, :]
batch.latents = latents
if batch.image_latent is not None:
image_latent = rearrange(batch.image_latent,
"b c (n t) h w -> b c n t h w",
n=sp_world_size).contiguous()
image_latent = image_latent[:, :, rank_in_sp_group, :, :, :]
batch.image_latent = image_latent
if not fastvideo_args.pipeline_config.ti2v_task and not fastvideo_args.pipeline_config.t2v_as_i2v_task:
image_latent = rearrange(batch.image_latent,
"b c (n t) h w -> b c n t h w",
n=sp_world_size).contiguous()
image_latent = image_latent[:, :, rank_in_sp_group, :, :, :]
batch.image_latent = image_latent
# Get timesteps and calculate warmup steps
timesteps = batch.timesteps
# TODO(will): remove this once we add input/output validation for stages
@@ -149,7 +158,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,15 +196,23 @@ 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:
boundary_timestep = fastvideo_args.boundary_ratio * self.scheduler.num_train_timesteps
if fastvideo_args.pipeline_config.dit_config.boundary_ratio is not None:
boundary_timestep = fastvideo_args.pipeline_config.dit_config.boundary_ratio
if batch.boundary_timestep is not None:
logger.info("Overriding boundary timestep from %s to %s",
boundary_timestep, batch.boundary_timestep)
boundary_timestep = batch.boundary_timestep
boundary_timestep *= self.scheduler.num_train_timesteps
else:
boundary_timestep = None
latent_model_input = latents.to(target_dtype)
@@ -236,6 +254,9 @@ class DenoisingStage(PipelineStage):
patch_size[2])
seq_len = int(math.ceil(seq_len / sp_world_size)) * sp_world_size
trajectory_timesteps: list[int] = []
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):
@@ -263,11 +284,27 @@ class DenoisingStage(PipelineStage):
# Expand latents for I2V
latent_model_input = latents.to(target_dtype)
if batch.image_latent is not None:
if batch.image_latent is not None and not fastvideo_args.pipeline_config.t2v_as_i2v_task:
assert not fastvideo_args.pipeline_config.ti2v_task, "image latents should not be provided for TI2V task"
latent_model_input = torch.cat(
[latent_model_input, batch.image_latent],
dim=1).to(target_dtype)
elif batch.image_latent is not None and fastvideo_args.pipeline_config.t2v_as_i2v_task:
assert batch.image_latent is not None, "image latents should be provided for T2V to I2V task"
if rank_in_sp_group == 0:
logger.info("latent_model_input.shape: %s",
latent_model_input.shape)
latent_model_input = torch.cat([
batch.image_latent,
latent_model_input[:, :, 1:, :, :],
],
dim=2).to(target_dtype)
logger.info("latent_model_input.shape: %s",
latent_model_input.shape)
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,9 +317,15 @@ 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)
if fastvideo_args.pipeline_config.t2v_as_i2v_task:
if rank_in_sp_group == 0:
latent_model_input = torch.cat([
batch.image_latent,
latent_model_input[:, :, 1:, :, :],
],
dim=2).to(target_dtype)
# Prepare inputs for transformer
guidance_expand = (
@@ -325,6 +368,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 +457,12 @@ 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.cpu())
trajectory_latents.append(latents)
# Update progress bar
if i == len(timesteps) - 1 or (
(i + 1) > num_warmup_steps and
@@ -397,8 +471,32 @@ class DenoisingStage(PipelineStage):
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)
else:
trajectory_tensor = None
if sp_group:
latents = sequence_model_parallel_all_gather(latents, dim=2)
if batch.return_trajectory_latents:
# logger.info("before stack trajectory_latents.shape: %s", trajectory_latents[0].shape)
logger.info("after stack trajectory_latents.shape: %s", trajectory_tensor.shape)
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:
batch.trajectory_timesteps = torch.tensor(trajectory_timesteps).cpu()
batch.trajectory_latents = trajectory_tensor.cpu()
if fastvideo_args.pipeline_config.t2v_as_i2v_task:
latents = torch.cat([
batch.image_latent,
latents[:, :, 1:, :, :],
],
dim=2)
# Update batch with final latents
batch.latents = latents
@@ -737,7 +835,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 +875,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])
+94 -48
View File
@@ -105,6 +105,81 @@ class ImageVAEEncodingStage(PipelineStage):
def __init__(self, vae: ParallelTiledVAE) -> None:
self.vae: ParallelTiledVAE = vae
def encode_image(self,
image: PIL.Image.Image,
height: int,
width: int,
fastvideo_args: FastVideoArgs,
generator: torch.Generator | None = None) -> torch.Tensor:
"""
Encode image into latent space.
"""
image = self.preprocess(
image,
vae_scale_factor=self.vae.spatial_compression_ratio,
height=height,
width=width).to(get_local_torch_device(), dtype=torch.float32)
# (B, C, H, W) -> (B, C, 1, H, W)
print(f"image.shape: {image.shape}")
image = image.unsqueeze(2)
print(f"after unsqueeze image.shape: {image.shape}")
return self.encode_tensor(image, fastvideo_args, generator)
def encode_tensor(self,
video_condition: torch.Tensor,
fastvideo_args: FastVideoArgs,
generator: torch.Generator | None = None) -> torch.Tensor:
"""
Encode frames into latent space.
"""
self.vae = self.vae.to(get_local_torch_device())
video_condition = video_condition.to(device=get_local_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)
if fastvideo_args.mode == ExecutionMode.PREPROCESS:
latent_condition = encoder_output.mean
else:
generator = 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
return latent_condition
def forward(
self,
batch: ForwardBatch,
@@ -157,58 +232,29 @@ class ImageVAEEncodingStage(PipelineStage):
# (B, C, H, W) -> (B, C, 1, H, W)
image = image.unsqueeze(2)
video_condition = torch.cat([
image,
image.new_zeros(image.shape[0], image.shape[1], num_frames - 1,
image.shape[3], image.shape[4])
],
dim=2)
video_condition = video_condition.to(device=get_local_torch_device(),
dtype=torch.float32)
if fastvideo_args.pipeline_config.t2v_as_i2v_task:
# repeat the image self.vae.temporal_compression_ratio times
video_condition = image.repeat(1, 1,
self.vae.temporal_compression_ratio,
1, 1)
# video_condition = image
logger.info("video_condition.shape: %s", video_condition.shape)
else:
video_condition = torch.cat([
image,
image.new_zeros(image.shape[0], image.shape[1], num_frames - 1,
image.shape[3], image.shape[4])
],
dim=2)
# 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)
if fastvideo_args.mode == ExecutionMode.PREPROCESS:
latent_condition = encoder_output.mean
else:
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
latent_condition = self.encode_tensor(video_condition, fastvideo_args,
batch.generator)
if fastvideo_args.mode == ExecutionMode.PREPROCESS:
batch.image_latent = latent_condition
elif fastvideo_args.pipeline_config.t2v_as_i2v_task:
logger.info("latent_condition.shape: %s", latent_condition.shape)
batch.image_latent = latent_condition
else:
mask_lat_size = torch.ones(1, 1, num_frames, latent_height,
latent_width)
@@ -35,9 +35,15 @@ class InputValidationStage(PipelineStage):
"""Generate seeds for the inference"""
seed = batch.seed
num_videos_per_prompt = batch.num_videos_per_prompt
if isinstance(batch.prompt, list):
num_prompts = len(batch.prompt)
else:
num_prompts = 1
total_num_videos = num_prompts * num_videos_per_prompt
assert seed is not None
seeds = [seed + i for i in range(num_videos_per_prompt)]
seeds = [seed + i for i in range(total_num_videos)]
batch.seeds = seeds
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
batch.generator = [
+2 -2
View File
@@ -82,8 +82,8 @@ def rocm_platform_plugin() -> str | None:
logger.info("ROCm platform is available")
finally:
amdsmi.amdsmi_shut_down()
except Exception as e:
logger.info("ROCm platform is unavailable: %s", e)
except Exception:
pass
return "fastvideo.platforms.rocm.RocmPlatform" if is_rocm else None
+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,123 @@
import numpy as np
from fastvideo.dataset.dataloader.record_schema import (
basic_t2v_record_creator,
i2v_record_creator,
ode_text_only_record_creator,
text_only_record_creator,
)
from fastvideo.pipelines.pipeline_batch_info import PreprocessBatch
def _mk_basic_batch(N: int) -> PreprocessBatch:
batch = PreprocessBatch(data_type="video")
batch.video_file_name = [f"vid_{i}" for i in range(N)]
batch.prompt = [f"caption_{i}" for i in range(N)]
batch.width = [640 for _ in range(N)]
batch.height = [360 for _ in range(N)]
batch.fps = [4 for _ in range(N)]
batch.num_frames = [2 for _ in range(N)]
# Latents: shape (N, C, T, H, W); per-record use latents[idx]
batch.latents = np.zeros((N, 4, 2, 8, 8), dtype=np.float32)
# Prompt embeds: list of per-record arrays [Seq, Dim]
batch.prompt_embeds = [np.ones((6, 16), dtype=np.float32) for _ in range(N)]
return batch
def test_basic_t2v_record_creator_fields():
N = 2
batch = _mk_basic_batch(N)
records = basic_t2v_record_creator(batch)
assert isinstance(records, list) and len(records) == N
for i, rec in enumerate(records):
assert rec["id"] == batch.video_file_name[i]
# Latents bytes/shape/dtype
assert isinstance(rec["vae_latent_bytes"], (bytes, bytearray))
assert rec["vae_latent_shape"] == list(batch.latents[i].shape)
assert rec["vae_latent_dtype"] == str(batch.latents[i].dtype)
# Text embedding
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
assert rec["text_embedding_shape"] == list(batch.prompt_embeds[i].shape)
assert rec["text_embedding_dtype"] == str(batch.prompt_embeds[i].dtype)
# Meta
assert rec["caption"] == batch.prompt[i]
assert rec["media_type"] == "video"
assert rec["width"] == int(batch.width[i])
assert rec["height"] == int(batch.height[i])
assert rec["num_frames"] == batch.latents[i].shape[1]
def test_i2v_record_creator_additional_fields():
N = 3
batch = _mk_basic_batch(N)
# image_embeds is a list of length 1, with an array of shape [N, D]
batch.image_embeds = [np.ones((N, 32), dtype=np.float32)]
# first frame latent per record
batch.image_latent = np.zeros((N, 4, 1, 8, 8), dtype=np.float32)
# pil image per record
batch.pil_image = np.zeros((N, 8, 8, 3), dtype=np.uint8)
records = i2v_record_creator(batch)
assert isinstance(records, list) and len(records) == N
for i, rec in enumerate(records):
# clip feature
assert isinstance(rec["clip_feature_bytes"], (bytes, bytearray))
assert rec["clip_feature_shape"] == list(batch.image_embeds[0][i].shape)
assert rec["clip_feature_dtype"] == str(batch.image_embeds[0][i].dtype)
# first frame latent
assert isinstance(rec["first_frame_latent_bytes"], (bytes, bytearray))
assert rec["first_frame_latent_shape"] == list(batch.image_latent[i].shape)
assert rec["first_frame_latent_dtype"] == str(batch.image_latent[i].dtype)
# pil image
assert isinstance(rec["pil_image_bytes"], (bytes, bytearray))
assert rec["pil_image_shape"] == list(batch.pil_image[i].shape)
assert rec["pil_image_dtype"] == str(batch.pil_image[i].dtype)
def test_ode_text_only_record_creator():
video_name = "ex"
caption = "a prompt"
text_embedding = np.ones((6, 16), dtype=np.float32)
traj = np.ones((5, 4, 2, 2), dtype=np.float32)
tsteps = np.arange(5, dtype=np.float32)
rec = ode_text_only_record_creator(
video_name=video_name,
text_embedding=text_embedding,
caption=caption,
trajectory_latents=traj,
trajectory_timesteps=tsteps,
)
assert rec["id"] == f"text_{video_name}"
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
assert rec["text_embedding_shape"] == list(text_embedding.shape)
assert rec["text_embedding_dtype"] == str(text_embedding.dtype)
assert rec["file_name"] == video_name
assert rec["caption"] == caption
assert rec["media_type"] == "text"
# Trajectory fields
assert isinstance(rec["trajectory_latents_bytes"], (bytes, bytearray))
assert rec["trajectory_latents_shape"] == list(traj.shape)
assert rec["trajectory_latents_dtype"] == str(traj.dtype)
assert isinstance(rec["trajectory_timesteps_bytes"], (bytes, bytearray))
assert rec["trajectory_timesteps_shape"] == list(tsteps.shape)
assert rec["trajectory_timesteps_dtype"] == str(tsteps.dtype)
def test_text_only_record_creator():
text_name = "note1"
caption = "a prompt"
text_embedding = np.ones((7, 16), dtype=np.float32)
rec = text_only_record_creator(
text_name=text_name,
text_embedding=text_embedding,
caption=caption,
)
assert rec["id"] == f"text_{text_name}"
assert isinstance(rec["text_embedding_bytes"], (bytes, bytearray))
assert rec["text_embedding_shape"] == list(text_embedding.shape)
assert rec["text_embedding_dtype"] == str(text_embedding.dtype)
assert rec["caption"] == caption
@@ -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()
+13 -1
View File
@@ -102,10 +102,22 @@ 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")
@app.function(gpu="L40S:1", image=image, timeout=900)
def run_unit_test():
run_test("pytest ./fastvideo/tests/dataset/ ./fastvideo/tests/workflow/ -vs")
@@ -0,0 +1,292 @@
# SPDX-License-Identifier: Apache-2.0
import os
import math
import numpy as np
import pytest
import torch
from fastvideo.wan.modules.causal_model import CausalWanModel
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.models.dits.causal_wanvideo import CausalWanTransformer3DModel
from fastvideo.utils import maybe_download_model
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
@pytest.mark.usefixtures("distributed_setup")
def test_ori_causal_wan_transformer():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
dit_cpu_offload=True,
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
model1 = CausalWanModel.from_pretrained(
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
causal_state_dict = torch.load("/mnt/weka/home/hao.zhang/wei/Self-Forcing/checkpoints/self_forcing_dmd.pt")["generator_ema"]
new_state_dict = {}
for k, v in causal_state_dict.items():
if k.startswith("model."):
new_state_dict[k.replace("model.", "")] = v
causal_state_dict = new_state_dict
model1.load_state_dict(causal_state_dict)
total_params = sum(p.numel() for p in model1.parameters())
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
weight_sum_model1 = sum(
p.to(torch.float64).sum().item() for p in model1.parameters())
# Also calculate mean for more stable comparison
weight_mean_model1 = weight_sum_model1 / total_params
logger.info("Model 1 weight sum: %s", weight_sum_model1)
logger.info("Model 1 weight mean: %s", weight_mean_model1)
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
total_params_model2 = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(
p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params_model2
logger.info("Model 2 weight sum: %s", weight_sum_model2)
logger.info("Model 2 weight mean: %s", weight_mean_model2)
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
logger.info("Weight sum difference: %s", weight_sum_diff)
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
logger.info("Weight mean difference: %s", weight_mean_diff)
# Set both models to eval mode
model1 = model1.eval()
model2 = model2.eval()
# Create identical inputs for both models
batch_size = 1
text_seq_len = 30
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(batch_size,
16,
12,
160,
90,
device=device,
dtype=precision)
block_sizes = [3 for _ in range(4)]
timesteps = [1000, 750, 500, 250]
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(batch_size,
text_seq_len + 1,
4096,
device=device,
dtype=precision)
output1 = _causal_inference(model1, hidden_states.clone(), encoder_hidden_states.clone(), block_sizes, timesteps, precision)
logger.info("Finish inference for model1")
output2 = _causal_inference(model2, hidden_states.clone(), encoder_hidden_states.clone(), block_sizes, timesteps, precision)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
logger.info("Output 1 Sum: %s", output1.float().sum().item())
logger.info("Output 2 Sum: %s", output2.float().sum().item())
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
def _causal_inference(transformer, latents, prompt_embeds, block_sizes, timesteps, target_dtype):
forward_batch = ForwardBatch(
data_type="dummy",
)
start_index = 0
pos_start_base = 0
frame_seq_length = latents.shape[-1] * latents.shape[-2] // (WanVideoConfig().arch_config.patch_size[-1] * WanVideoConfig().arch_config.patch_size[-2])
seq_len = frame_seq_length * latents.shape[2]
kv_cache1 = _initialize_kv_cache(transformer, batch_size=latents.shape[0],
kv_cache_size=frame_seq_length * latents.shape[2],
dtype=target_dtype,
device=latents.device)
crossattn_cache = _initialize_crossattn_cache(
transformer,
batch_size=latents.shape[0],
max_text_len=WanVideoConfig().arch_config.text_len,
dtype=target_dtype,
device=latents.device)
for current_num_frames, t_cur in zip(block_sizes, timesteps):
# logger.info(f"Current frame idx: {start_index}, Current timestep: {t_cur}")
# logger.info(f"k cache sum: {sum(kv_cache['k'].float().sum().item() for kv_cache in kv_cache1)}, v cache sum: {sum(kv_cache['v'].float().sum().item() for kv_cache in kv_cache1)}")
# logger.info(f"latents sum: {latents.float().sum().item()}, encoder_hidden_states sum: {prompt_embeds.float().sum().item()}")
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
attn_metadata = None
with set_forward_context(current_timestep=0,
attn_metadata=attn_metadata,
forward_batch=forward_batch):
# Run transformer; follow DMD stage pattern
t_expanded_noise = t_cur * torch.ones(
(current_latents.shape[0], 1),
device=current_latents.device,
dtype=torch.long)
if isinstance(transformer, CausalWanModel):
pred_noise_btchw = transformer(
x=current_latents,
context=prompt_embeds,
t=t_expanded_noise,
seq_len=seq_len,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
frame_seq_length
)
elif isinstance(transformer, CausalWanTransformer3DModel):
pred_noise_btchw = transformer(
current_latents,
prompt_embeds,
t_expanded_noise,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
frame_seq_length,
start_frame=start_index
)
# Write back and advance
latents[:, :, start_index:start_index +
current_num_frames, :, :] = pred_noise_btchw.clone()
# Re-run with context timestep to update KV cache using clean context
context_noise = 0
t_context = torch.ones([latents.shape[0]],
device=latents.device,
dtype=torch.long) * int(context_noise)
context_bcthw = pred_noise_btchw.to(target_dtype)
with set_forward_context(current_timestep=0,
attn_metadata=attn_metadata,
forward_batch=forward_batch):
t_expanded_context = t_context.unsqueeze(1)
if isinstance(transformer, CausalWanModel):
_ = transformer(
x=context_bcthw,
context=prompt_embeds,
t=t_expanded_context,
seq_len=seq_len,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
frame_seq_length
)
elif isinstance(transformer, CausalWanTransformer3DModel):
_ = transformer(
context_bcthw,
prompt_embeds,
t_expanded_context,
kv_cache=kv_cache1,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
frame_seq_length,
start_frame=start_index
)
start_index += current_num_frames
return latents
def _initialize_kv_cache(transformer, batch_size, kv_cache_size, dtype, device) -> None:
"""
Initialize a Per-GPU KV cache aligned with the Wan model assumptions.
"""
kv_cache1 = []
if isinstance(transformer, CausalWanModel):
num_attention_heads = transformer.num_heads
attention_head_dim = transformer.dim // transformer.num_heads
elif isinstance(transformer, CausalWanTransformer3DModel):
num_attention_heads = transformer.num_attention_heads
attention_head_dim = transformer.attention_head_dim
for _ in range(len(transformer.blocks)):
kv_cache1.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),
})
return kv_cache1
def _initialize_crossattn_cache(transformer, batch_size, max_text_len, dtype,
device) -> None:
"""
Initialize a Per-GPU cross-attention cache aligned with the Wan model assumptions.
"""
crossattn_cache = []
if isinstance(transformer, CausalWanModel):
num_attention_heads = transformer.num_heads
attention_head_dim = transformer.dim // transformer.num_heads
elif isinstance(transformer, CausalWanTransformer3DModel):
num_attention_heads = transformer.num_attention_heads
attention_head_dim = transformer.attention_head_dim
for _ in range(len(transformer.blocks)):
crossattn_cache.append({
"k":
torch.zeros([
batch_size, max_text_len, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"v":
torch.zeros([
batch_size, max_text_len, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"is_init":
False,
})
return crossattn_cache
@@ -0,0 +1,133 @@
# SPDX-License-Identifier: Apache-2.0
import os
import math
import numpy as np
import pytest
import torch
from fastvideo.wan.modules.model import WanModel
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.utils import maybe_download_model
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
@pytest.mark.usefixtures("distributed_setup")
def test_ori_wan_transformer():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
dit_cpu_offload=True,
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
model1 = WanModel.from_pretrained(
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
total_params = sum(p.numel() for p in model1.parameters())
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
weight_sum_model1 = sum(
p.to(torch.float64).sum().item() for p in model1.parameters())
# Also calculate mean for more stable comparison
weight_mean_model1 = weight_sum_model1 / total_params
logger.info("Model 1 weight sum: %s", weight_sum_model1)
logger.info("Model 1 weight mean: %s", weight_mean_model1)
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
total_params_model2 = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(
p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params_model2
logger.info("Model 2 weight sum: %s", weight_sum_model2)
logger.info("Model 2 weight mean: %s", weight_mean_model2)
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
logger.info("Weight sum difference: %s", weight_sum_diff)
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
logger.info("Weight mean difference: %s", weight_mean_diff)
# Set both models to eval mode
model1 = model1.eval()
model2 = model2.eval()
# Create identical inputs for both models
batch_size = 1
text_seq_len = 30
seq_len = math.ceil((160 * 90) /
(2 * 2) *
21)
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(batch_size,
16,
21,
160,
90,
device=device,
dtype=precision)
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(batch_size,
text_seq_len + 1,
4096,
device=device,
dtype=precision)
# Timestep
timestep = torch.tensor([500], device=device, dtype=precision)
forward_batch = ForwardBatch(
data_type="dummy",
)
# with torch.amp.autocast('cuda', dtype=precision):
output1 = model1(
x=hidden_states,
context=encoder_hidden_states,
t=timestep,
seq_len=seq_len,
)
with set_forward_context(
current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch,
):
output2 = model2(hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
@@ -0,0 +1,144 @@
# SPDX-License-Identifier: Apache-2.0
import os
import math
import numpy as np
import pytest
import torch
from fastvideo.wan.modules.causal_model import CausalWanModel
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TransformerLoader
from fastvideo.utils import maybe_download_model
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
BASE_MODEL_PATH = "wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH = maybe_download_model(BASE_MODEL_PATH,
local_dir=os.path.join(
'data', BASE_MODEL_PATH))
TRANSFORMER_PATH = os.path.join(MODEL_PATH, "transformer")
@pytest.mark.usefixtures("distributed_setup")
def test_train_ori_causal_wan_transformer():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision = torch.bfloat16
precision_str = "bf16"
args = FastVideoArgs(model_path=TRANSFORMER_PATH,
dit_cpu_offload=True,
pipeline_config=PipelineConfig(dit_config=WanVideoConfig(), dit_precision=precision_str))
args.device = device
loader = TransformerLoader()
model2 = loader.load(TRANSFORMER_PATH, args).to(dtype=precision)
model1 = CausalWanModel.from_pretrained(
"/mnt/weka/home/hao.zhang/wei/Self-Forcing/wan_models/Wan2.1-T2V-1.3B", device=device,
torch_dtype=precision).to(device, dtype=precision).requires_grad_(False)
causal_state_dict = torch.load("/mnt/weka/home/hao.zhang/wei/Self-Forcing/checkpoints/self_forcing_dmd.pt")["generator_ema"]
new_state_dict = {}
for k, v in causal_state_dict.items():
if k.startswith("model."):
new_state_dict[k.replace("model.", "")] = v
causal_state_dict = new_state_dict
model1.load_state_dict(causal_state_dict)
model1.num_frame_per_block = 3
model2.num_frame_per_block = 3
total_params = sum(p.numel() for p in model1.parameters())
# Calculate weight sum for model1 (converting to float64 to avoid overflow)
weight_sum_model1 = sum(
p.to(torch.float64).sum().item() for p in model1.parameters())
# Also calculate mean for more stable comparison
weight_mean_model1 = weight_sum_model1 / total_params
logger.info("Model 1 weight sum: %s", weight_sum_model1)
logger.info("Model 1 weight mean: %s", weight_mean_model1)
# Calculate weight sum for model2 (converting to float64 to avoid overflow)
total_params_model2 = sum(p.numel() for p in model2.parameters())
weight_sum_model2 = sum(
p.to(torch.float64).sum().item() for p in model2.parameters())
# Also calculate mean for more stable comparison
weight_mean_model2 = weight_sum_model2 / total_params_model2
logger.info("Model 2 weight sum: %s", weight_sum_model2)
logger.info("Model 2 weight mean: %s", weight_mean_model2)
weight_sum_diff = abs(weight_sum_model1 - weight_sum_model2)
logger.info("Weight sum difference: %s", weight_sum_diff)
weight_mean_diff = abs(weight_mean_model1 - weight_mean_model2)
logger.info("Weight mean difference: %s", weight_mean_diff)
# Set both models to eval mode
model1 = model1.eval()
model2 = model2.eval()
# Create identical inputs for both models
batch_size = 1
text_seq_len = 30
seq_len = math.ceil((160 * 90) /
(2 * 2) *
21)
# Video latents [B, C, T, H, W]
hidden_states = torch.randn(batch_size,
16,
21,
160,
90,
device=device,
dtype=precision)
# Text embeddings [B, L, D] (including global token)
encoder_hidden_states = torch.randn(batch_size,
text_seq_len + 1,
4096,
device=device,
dtype=precision)
# Timestep
timestep = torch.randint(0, 1000, (batch_size, 21), device=device, dtype=torch.long)
logger.info("timestep: %s", timestep)
forward_batch = ForwardBatch(
data_type="dummy",
)
# with torch.amp.autocast('cuda', dtype=precision):
output1 = model1(
x=hidden_states,
context=encoder_hidden_states,
t=timestep,
seq_len=seq_len,
)
with set_forward_context(
current_timestep=0,
attn_metadata=None,
forward_batch=forward_batch,
):
output2 = model2(hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=timestep)
# Check if outputs have the same shape
assert output1.shape == output2.shape, f"Output shapes don't match: {output1.shape} vs {output2.shape}"
assert output1.dtype == output2.dtype, f"Output dtype don't match: {output1.dtype} vs {output2.dtype}"
# Check if outputs are similar (allowing for small numerical differences)
max_diff = torch.max(torch.abs(output1 - output2))
mean_diff = torch.mean(torch.abs(output1 - output2))
logger.info("Max Diff: %s", max_diff.item())
logger.info("Mean Diff: %s", mean_diff.item())
assert max_diff < 1e-4, f"Maximum difference between outputs: {max_diff.item()}"
# mean diff
assert mean_diff < 1e-4, f"Mean difference between outputs: {mean_diff.item()}"
@@ -0,0 +1,68 @@
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)],
# Attention mask should be integer dtype in real pipelines
prompt_attention_mask=[torch.ones(B, 1, dtype=torch.int64)],
)
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()
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_tables(write_remainder=True)
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
+407 -73
View File
@@ -1,6 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import copy
import gc
import json
import os
import time
from abc import abstractmethod
@@ -11,6 +12,7 @@ from typing import Any
import imageio
import numpy as np
import torch
import torch.distributed as dist
import torch.nn.functional as F
import torchvision
from einops import rearrange
@@ -36,9 +38,11 @@ from fastvideo.training.activation_checkpoint import (
apply_activation_checkpointing)
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases, get_scheduler,
load_distillation_checkpoint, save_distillation_checkpoint, shift_timestep)
from fastvideo.utils import is_vsa_available, set_random_seed
EMA_FSDP, clip_grad_norm_while_handling_failing_dtensor_cases,
get_scheduler, load_distillation_checkpoint, save_distillation_checkpoint,
shift_timestep)
from fastvideo.utils import (is_vsa_available, maybe_download_model,
set_random_seed, verify_model_config_and_directory)
import wandb # isort: skip
@@ -87,9 +91,27 @@ class DistillationPipeline(TrainingPipeline):
self.noise_scheduler = FlowMatchEulerDiscreteScheduler(
shift=self.timestep_shift)
# 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")
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.real_score_transformer.eval()
@@ -116,10 +138,13 @@ 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=(0.9, 0.999),
betas=betas,
weight_decay=training_args.weight_decay,
eps=1e-8,
)
@@ -147,8 +172,19 @@ class DistillationPipeline(TrainingPipeline):
self.training_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long,
device=get_local_torch_device())
logger.info("Distillation generator model to %s denoising steps",
len(self.denoising_step_list))
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)
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
self.min_timestep = int(self.training_args.min_timestep_ratio *
@@ -158,6 +194,82 @@ 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."""
@@ -174,6 +286,110 @@ class DistillationPipeline(TrainingPipeline):
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,
@@ -330,6 +546,7 @@ class DistillationPipeline(TrainingPipeline):
def _dmd_forward(self, generator_pred_video: torch.Tensor,
training_batch: TrainingBatch) -> torch.Tensor:
"""Compute DMD (Diffusion Model Distillation) loss."""
original_latent = generator_pred_video
with torch.no_grad():
timestep = torch.randint(0,
self.num_train_timestep, [1],
@@ -354,7 +571,7 @@ class DistillationPipeline(TrainingPipeline):
noisy_latent = self.noise_scheduler.add_noise(
generator_pred_video.flatten(0, 1), noise.flatten(0, 1),
timestep).unflatten(0, (1, generator_pred_video.shape[1]))
timestep).detach().unflatten(0, (1, generator_pred_video.shape[1]))
# fake_score_transformer forward
training_batch = self._build_distill_input_kwargs(
@@ -403,24 +620,24 @@ class DistillationPipeline(TrainingPipeline):
pred_real_video_uncond) * self.real_score_guidance_scale
grad = (faker_score_pred_video - real_score_pred_video) / torch.abs(
generator_pred_video - real_score_pred_video).mean()
original_latent - real_score_pred_video).mean()
grad = torch.nan_to_num(grad)
dmd_loss = 0.5 * F.mse_loss(
generator_pred_video.float(),
(generator_pred_video.float() - grad.float()).detach())
original_latent.float(),
(original_latent.float() - grad.float()).detach())
training_batch.dmd_latent_vis_dict.update({
"training_batch_dmd_fwd_clean_latent":
training_batch.latents,
"generator_pred_video":
generator_pred_video,
original_latent.detach(),
"real_score_pred_video":
real_score_pred_video,
real_score_pred_video.detach(),
"faker_score_pred_video":
faker_score_pred_video,
faker_score_pred_video.detach(),
"dmd_timestep":
timestep,
timestep.detach(),
})
return dmd_loss
@@ -517,12 +734,12 @@ class DistillationPipeline(TrainingPipeline):
"encoder_hidden_states": self.negative_prompt_embeds,
"encoder_attention_mask": self.negative_prompt_attention_mask,
}
training_batch.unconditional_dict = unconditional_dict
training_batch.dmd_latent_vis_dict = {}
training_batch.fake_score_latent_vis_dict = {}
training_batch.conditional_dict = conditional_dict
training_batch.unconditional_dict = unconditional_dict
training_batch.raw_latent_shape = training_batch.latents.shape
training_batch.latents = training_batch.latents.permute(0, 2, 1, 3, 4)
self.video_latent_shape = training_batch.latents.shape
@@ -585,8 +802,15 @@ class DistillationPipeline(TrainingPipeline):
(dmd_loss / gradient_accumulation_steps).backward()
total_dmd_loss += dmd_loss.detach().item()
self._clip_model_grad_norm_(batch_gen, self.transformer)
for param in self.transformer.parameters():
# check if the gradient is not None and not zero
assert param.grad is not None and param.grad.abs().sum() > 0
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)
@@ -610,6 +834,9 @@ class DistillationPipeline(TrainingPipeline):
fake_score_latent_vis_dict.update(
batch_fake.fake_score_latent_vis_dict)
self._clip_model_grad_norm_(batch_fake, self.fake_score_transformer)
for param in self.fake_score_transformer.parameters():
# check if the gradient is not None and not zero
assert param.grad is not None and param.grad.abs().sum() > 0
self.fake_score_optimizer.step()
self.fake_score_lr_scheduler.step()
self.lr_scheduler.step()
@@ -637,7 +864,8 @@ 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.fake_score_lr_scheduler, self.noise_random_generator,
self.generator_ema)
if resumed_step > 0:
self.init_steps = resumed_step
@@ -668,6 +896,14 @@ 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
@@ -699,6 +935,18 @@ 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]
@@ -714,50 +962,98 @@ class DistillationPipeline(TrainingPipeline):
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)
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)
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)
# 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)
# Log validation results for this step
world_group = get_world_group()
@@ -834,16 +1130,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:
@@ -904,6 +1200,10 @@ 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()
@@ -947,6 +1247,14 @@ 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)
@@ -960,11 +1268,19 @@ 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,
"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)
@@ -992,6 +1308,15 @@ 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":
@@ -1023,7 +1348,8 @@ 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.fake_score_lr_scheduler, self.noise_random_generator,
self.generator_ema)
if self.transformer:
self.transformer.train()
@@ -1040,7 +1366,11 @@ class DistillationPipeline(TrainingPipeline):
self.global_rank,
self.training_args.output_dir,
f"{step}_weight_only",
only_save_generator_weight=True)
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:
if self.training_args.log_visualization:
@@ -1060,7 +1390,11 @@ 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.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()
cleanup_dist_env_and_memory()
+443
View File
@@ -0,0 +1,443 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from typing import cast
import torch
import torch.nn.functional as F
import numpy as np
import wandb
from fastvideo.dataset.dataloader.schema import pyarrow_schema_ode_trajectory_text_only
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
SelfForcingFlowMatchScheduler)
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
WanCausalDMDPipeline)
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
from fastvideo.training.training_utils import (
clip_grad_norm_while_handling_failing_dtensor_cases)
logger = init_logger(__name__)
class ODEInitTrainingPipeline(TrainingPipeline):
"""
Training pipeline for ODE-init using precomputed denoising trajectories.
Supervision: predict the next latent in the stored trajectory by
- feeding current latent at timestep t into the transformer to predict noise
- stepping the scheduler with the predicted noise
- minimizing MSE to the stored next latent at timestep t_next
"""
_required_config_modules = ["scheduler", "transformer", "vae"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# Match the preprocess/generation scheduler for consistent stepping
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=1000,
training=True)
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_ode_trajectory_text_only
def initialize_training_pipeline(self, training_args: TrainingArgs):
super().initialize_training_pipeline(training_args)
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
assert self.timestep_shift == 5.0, "flow_shift must be 5.0"
self.noise_scheduler = SelfForcingFlowMatchScheduler(
shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
self.noise_scheduler.set_timesteps(num_inference_steps=1000,
training=True)
# logger.info(f"ARG dmd_denoising_steps: {training_args.pipeline_config.dmd_denoising_steps}")
logger.info(
f"ARG dmd_denoising_steps: {self.training_args.pipeline_config.dmd_denoising_steps}"
)
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250],
dtype=torch.long,
device=get_local_torch_device())
# self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250], 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()
logger.info(f"timesteps: {timesteps}")
self.dmd_denoising_steps = timesteps[1000 -
self.dmd_denoising_steps]
logger.info(
f"warped self.dmd_denoising_steps: {self.dmd_denoising_steps}")
# assert False, "warp_denoising_step must be false"
else:
assert False, "warp_denoising_step must be true"
logger.info("not warped")
self.dmd_denoising_steps = self.dmd_denoising_steps.to(
get_local_torch_device())
logger.info(f"denoising_step_list: {self.dmd_denoising_steps}")
logger.info(
"Initialized ODE-init training pipeline with %s denoising steps",
len(self.dmd_denoising_steps))
# Cache for nearest trajectory index per DMD step (computed lazily on first batch)
self._cached_closest_idx_per_dmd = None
self.num_train_timestep = self.noise_scheduler.num_train_timesteps
# self.min_timestep = int(self.training_args.min_timestep_ratio *
# self.num_train_timestep)
# self.max_timestep = int(self.training_args.max_timestep_ratio *
# self.num_train_timestep)
# self.real_score_guidance_scale = self.training_args.real_score_guidance_scale
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
# Warm start validation with current transformer
self.validation_pipeline = WanCausalDMDPipeline.from_pretrained(
# training_args.model_path,
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers",
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)
def _get_next_batch(self, training_batch): # type: ignore[override]
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)
self.train_loader_iter = iter(self.train_dataloader)
batch = next(self.train_loader_iter)
# Required fields from parquet (ODE trajectory schema)
encoder_hidden_states = batch['text_embedding']
encoder_attention_mask = batch['text_attention_mask']
infos = batch['info_list']
# Trajectory tensors may include a leading singleton batch dim per row
trajectory_latents = batch['trajectory_latents']
if trajectory_latents.dim() == 7:
# [B, 1, S, C, T, H, W] -> [B, S, C, T, H, W]
trajectory_latents = trajectory_latents[:, 0]
elif trajectory_latents.dim() == 6:
# already [B, S, C, T, H, W]
pass
else:
raise ValueError(
f"Unexpected trajectory_latents dim: {trajectory_latents.dim()}"
)
trajectory_timesteps = batch['trajectory_timesteps']
if trajectory_timesteps.dim() == 3:
# [B, 1, S] -> [B, S]
trajectory_timesteps = trajectory_timesteps[:, 0]
elif trajectory_timesteps.dim() == 2:
# [B, S]
pass
else:
raise ValueError(
f"Unexpected trajectory_timesteps dim: {trajectory_timesteps.dim()}"
)
# [B, S, C, T, H, W] -> [B, S, T, C, H, W] to match self-forcing
trajectory_latents = trajectory_latents.permute(0, 1, 3, 2, 4, 5)
# Move to device
device = get_local_torch_device()
training_batch.encoder_hidden_states = encoder_hidden_states.to(
device, dtype=torch.bfloat16)
training_batch.encoder_attention_mask = encoder_attention_mask.to(
device, dtype=torch.bfloat16)
training_batch.infos = infos
return training_batch, trajectory_latents.to(
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
def _get_timestep(self,
min_timestep: int,
max_timestep: int,
batch_size: int,
num_frame: int,
num_frame_per_block: int,
uniform_timestep: bool = False) -> torch.Tensor:
if uniform_timestep:
timestep = torch.randint(min_timestep,
max_timestep, [batch_size, 1],
device=self.device,
dtype=torch.long).repeat(1, num_frame)
return timestep
else:
timestep = torch.randint(min_timestep,
max_timestep, [batch_size, num_frame],
device=self.device,
dtype=torch.long)
# logger.info(f"individual timestep: {timestep}")
# make the noise level the same within every block
timestep = timestep.reshape(timestep.shape[0], -1,
num_frame_per_block)
timestep[:, :, 1:] = timestep[:, :, 0:1]
timestep = timestep.reshape(timestep.shape[0], -1)
return timestep
def _step_predict_next_latent(
self, traj_latents: torch.Tensor, traj_timesteps: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_attention_mask: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, dict[str, torch.Tensor]]:
latent_vis_dict = {}
device = get_local_torch_device()
target_latent = traj_latents[:, -1]
# logger.info(f"traj_latents: {traj_latents.shape}")
# logger.info(f"traj_timesteps: {traj_timesteps.shape}")
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
B, S, num_frames, num_channels, height, width = traj_latents.shape
# Lazily cache nearest trajectory index per DMD step based on the (fixed) S timesteps
if self._cached_closest_idx_per_dmd is None:
# Use the first sample's trajectory timesteps; assumed identical across batches
# s_steps = traj_timesteps[0].to(torch.long) # [S]
# dmd = cast(torch.Tensor, self.dmd_denoising_steps).to(s_steps.device) # [K]
# distances_ks: [K, S] = |s_steps - dmd|
# distances_ks = (s_steps.unsqueeze(0) - dmd.unsqueeze(1)).abs()
# self._cached_closest_idx_per_dmd = distances_ks.argmin(dim=1).to(torch.long).cpu() # [K]
self._cached_closest_idx_per_dmd = torch.tensor(
[0, 12, 24, 36], dtype=torch.long).cpu()
logger.info(
f"self._cached_closest_idx_per_dmd: {self._cached_closest_idx_per_dmd}"
)
logger.info(
f"corresponding timesteps: {self.noise_scheduler.timesteps[self._cached_closest_idx_per_dmd]}"
)
# logger.info(f"traj_latents: {traj_latents.shape}")
# Select the K indexes from traj_latents using self._cached_closest_idx_per_dmd
# traj_latents: [B, S, C, T, H, W], self._cached_closest_idx_per_dmd: [K]
# Output: [B, K, C, T, H, W]
relevant_traj_latents = torch.index_select(
traj_latents,
dim=1,
index=self._cached_closest_idx_per_dmd.to(traj_latents.device))
# assert relevant_traj_latents.shape[0] == 1
indexes = self._get_timestep( # [B, num_frames]
0,
len(self.dmd_denoising_steps),
B,
num_frames,
3,
uniform_timestep=False)
logger.info(f"indexes: {indexes.shape}")
logger.info(f"indexes: {indexes}")
# noisy_input = relevant_traj_latents[indexes]
noisy_input = torch.gather(
relevant_traj_latents,
dim=1,
index=indexes.reshape(B, 1, num_frames, 1, 1,
1).expand(-1, -1, -1, num_channels, height,
width).to(self.device)).squeeze(1)
# noisy_input = noisy_input.unsqueeze(0)
# # Sample a single DMD step for the whole batch and fetch its cached nearest S-index
# K = len(self.dmd_denoising_steps)
# dmd_idx = torch.randint(0, K, (1,), device=device)
# logger.info(f"dmd_idx: {dmd_idx}")
# assert self._cached_closest_idx_per_dmd is not None
# nearest_s_idx = int(self._cached_closest_idx_per_dmd[int(dmd_idx.item())])
# nearest_idx = torch.full((B,), nearest_s_idx, device=device, dtype=torch.long)
# batch_indices = torch.arange(B, device=device)
# noisy_input = traj_latents[batch_indices, nearest_idx] # [B, C, T, H, W]
# target_latent = traj_latents[batch_indices, -1] # [B, C, T, H, W]
# t = traj_timesteps[batch_indices, nearest_idx] # [B]
# Scale model input as in inference for consistency with stored trajectories
# noisy_input = self.modules["scheduler"].scale_model_input(noisy_input, t)
# logger.info(f"indexes: {indexes.shape}")
# logger.info(f"indexes: {indexes}")
timestep = self.dmd_denoising_steps[indexes]
# logger.info(f"timestep: {timestep.shape}")
# logger.info(f"timestep: {timestep}")
# Prepare inputs for transformer
latent_vis_dict["noisy_input"] = noisy_input.permute(0, 2, 1, 3, 4).detach().clone().cpu()
latent_vis_dict["x0"] = target_latent.permute(0, 2, 1, 3, 4).detach().clone().cpu()
model_dtype = next(self.transformer.parameters()).dtype
input_kwargs = {
"hidden_states": noisy_input.permute(0, 2, 1, 3, 4),
"encoder_hidden_states": encoder_hidden_states,
"timestep": timestep.to(device, dtype=model_dtype),
"encoder_attention_mask": encoder_attention_mask,
"return_dict": False,
}
# Predict noise and step the scheduler to obtain next latent
with set_forward_context(current_timestep=timestep,
attn_metadata=None,
forward_batch=None):
noise_pred = self.transformer(**input_kwargs).permute(0, 2, 1, 3, 4)
# logger.info(f"noise_pred: {noise_pred.shape}")
if isinstance(noise_pred, (tuple, list)):
noise_pred = noise_pred[0]
from fastvideo.models.utils import pred_noise_to_pred_video
pred_video = pred_noise_to_pred_video(
pred_noise=noise_pred.flatten(0, 1),
noise_input_latent=noisy_input.flatten(0, 1),
timestep=timestep.to(dtype=model_dtype).flatten(0, 1),
scheduler=self.modules["scheduler"]).unflatten(
0, noise_pred.shape[:2])
latent_vis_dict["pred_video"] = pred_video.permute(0, 2, 1, 3, 4).detach().clone().cpu()
# noisy_input = pred_noise_to_pred_video(noise_pred, noisy_input, t, self.modules["scheduler"])
# next_latent_pred = self.modules["scheduler"].step(
# noise_pred, t, current_latents, return_dict=False)[0]
return pred_video, target_latent, timestep, latent_vis_dict
def train_one_step(self, training_batch): # type: ignore[override]
self.transformer.train()
self.optimizer.zero_grad()
training_batch.total_loss = 0.0
args = cast(TrainingArgs, self.training_args)
# Using cached nearest index per DMD step; computation happens in _step_predict_next_latent
for _ in range(args.gradient_accumulation_steps):
training_batch, traj_latents, traj_timesteps = self._get_next_batch(
training_batch)
text_embeds = training_batch.encoder_hidden_states
text_attention_mask = training_batch.encoder_attention_mask
assert traj_latents.shape[0] == 1
# Shapes: traj_latents [B, S, C, T, H, W], traj_timesteps [B, S]
B, S = traj_latents.shape[0], traj_latents.shape[1]
if S < 2:
raise ValueError("Trajectory must contain at least 2 steps")
# Sample per-sample current step i in [0, S-2]
# idx = torch.randint(low=0, high=S - 1, size=(B, ),
# device=traj_latents.device)
# Gather current latents and next latents
# batch_indices = torch.arange(B, device=traj_latents.device)
# current_latents = traj_latents[batch_indices, idx] # [B, C, T,H,W]
# current_latent = traj_timesteps[:, -1, :, :, :, :]
# target_latents = traj_latents[:, -1, :, :, :, :]
# Corresponding timesteps t (long) -> cast per sample
# t = traj_timesteps[:, -1, :, :, :, :]
# if t.dtype != torch.long:
# t = t.long()
# Forward to predict next latent by stepping scheduler with predicted noise
noise_pred, target_latent, t, latent_vis_dict = self._step_predict_next_latent(
traj_latents, traj_timesteps, text_embeds, text_attention_mask)
training_batch.latent_vis_dict.update(latent_vis_dict)
mask = t != 0
# Compute loss
loss = F.mse_loss(noise_pred[mask],
target_latent[mask],
reduction="mean")
loss = loss / args.gradient_accumulation_steps
with set_forward_context(current_timestep=t,
attn_metadata=None,
forward_batch=None):
loss.backward()
avg_loss = loss.detach().clone()
training_batch.total_loss += avg_loss.item()
# Clip grad and step optimizers
grad_norm = clip_grad_norm_while_handling_failing_dtensor_cases(
[p for p in self.transformer.parameters() if p.requires_grad],
args.max_grad_norm if args.max_grad_norm is not None else 0.0)
self.optimizer.step()
self.lr_scheduler.step()
if grad_norm is None:
grad_value = 0.0
else:
try:
if isinstance(grad_norm, torch.Tensor):
grad_value = float(grad_norm.detach().float().item())
else:
grad_value = float(grad_norm)
except Exception:
grad_value = 0.0
training_batch.grad_norm = grad_value
return training_batch
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 = {}
latents_vis_dict = training_batch.latent_vis_dict
latent_log_keys = ['noisy_input', 'x0', 'pred_video']
for latent_key in latent_log_keys:
assert latent_key in latents_vis_dict and latents_vis_dict[latent_key] is not None
latent = latents_vis_dict[latent_key]
pixel_latent = self.validation_pipeline.decoding_stage.decode(latent, training_args)
video = pixel_latent.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=16, format="mp4") # change to 16 for Wan2.1
# Clean up references
del video, pixel_latent, latent
# Log to wandb
if self.global_rank == 0:
wandb.log(wandb_loss_dict, step=step)
# dmd_latents_vis_dict = training_batch.dmd_latent_vis_dict
# fake_score_latents_vis_dict = training_batch.fake_score_latent_vis_dict
# fake_score_log_keys = ['generator_pred_video']
# dmd_log_keys = ['faker_score_pred_video', 'real_score_pred_video']
def main(args) -> None:
logger.info("Starting ODE-init training pipeline...")
logger.info(f"ARG dmd_denoising_steps: {args.dmd_denoising_steps}")
pipeline = ODEInitTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("ODE-init training pipeline done")
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()
args.dit_cpu_offload = False
main(args)

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