Compare commits

...
Author SHA1 Message Date
Wei Zhou ec199c8b2f Update preprocess_pipeline_ode_trajectory.py 2025-12-22 02:56:41 -08:00
JerryZhou54 0dc0617ac5 ckpt 2025-12-22 10:50:51 +00:00
JerryZhou54 087f5b927f Hardcode ode trajectory generation for hy1.5 2025-12-22 09:07:20 +00:00
JerryZhou54 9d822d464b lint 2025-12-21 05:38:44 +00:00
JerryZhou54 9f86b28ecb Add 720p config & support paddded SP 2025-12-21 05:37:08 +00:00
JerryZhou54 95b9cde729 checkpoint 2025-12-21 05:37:08 +00:00
JerryZhou54 54b85d8931 Fix lint 2025-12-21 05:37:08 +00:00
JerryZhou54 922b082cb2 Remove tracked files: test_ori_hunyuanvideo15.py 2025-12-21 05:37:08 +00:00
JerryZhou54 fb3bfecd18 Pipeline working 2025-12-21 05:37:06 +00:00
JerryZhou54 f01fb88c6a ckpt 2025-12-21 05:32:26 +00:00
JerryZhou54 e8f298c1ab 0 numerical diff between original and fv's hy1.5 dit 2025-12-21 05:32:26 +00:00
JerryZhou54 d7c7d23375 0 numerical diff between hy1.5 dit and original's 2025-12-21 05:32:26 +00:00
JerryZhou54 44808ce145 Added hy1.5 vae with 0 numerical diff with diffusers 2025-12-21 05:32:26 +00:00
JerryZhou54 0d5124e092 Qwen2_5_VL text encoder has 0 numerical diff with transformers 2025-12-21 05:32:26 +00:00
William Lin 7f71994653 [misc] Allow manual override of Pipeline class through override_pipeline_cls_name (#945) 2025-12-20 14:39:17 -06:00
Loay Rashid 2bb3349da1 [bugfix] Added VSA Padding logic (#944) 2025-12-20 14:29:11 -06:00
Kaiqin Kong 8fe1689968 [feat] Add Matrix-Game 2.0 (#938) 2025-12-20 14:09:12 -06:00
Loay Rashid e53730f324 [docs] Minor Fixes (#942) 2025-12-19 16:48:16 -06:00
Loay Rashid 7a4fe9086a [feat] Support sequence packing and shard after pachification for USP (#894) 2025-12-19 16:19:46 -06:00
Ohm-Rishabh d277361aae [misc] add schedule configurations to pytorch profiler (#934) 2025-12-18 01:45:23 -06:00
alexzms 734a54e7a9 [ci]: Use pre-built docker image & skip VSA compilation (#939) 2025-12-16 23:14:11 -08:00
alexzms 91364982df [Feature] Support for Variable Q/KV Sequence Lengths in VSA ThunderKittens kernel (#911) 2025-12-16 20:08:15 -08:00
William Lin 50145e4fcb [CI] Fix CI tests (#935) 2025-12-16 04:59:43 -08:00
William Lin 4112507e99 [misc] upgrade pytorch version to 2.9.0 (#928) 2025-12-15 04:12:43 -08:00
William Lin 424fc2b4ae [bugfix] [lora] [distillation] Fix lora distillation bug (#933) 2025-12-15 04:12:02 -08:00
William Lin e6066223e6 [bugfix] [VSA] [distillation] Various bugfixes for VSA and distillation and nightly tests (#932) 2025-12-12 16:51:54 -08:00
120 changed files with 8965 additions and 957 deletions
+2 -2
View File
@@ -15,7 +15,7 @@ FastVideo features an end-to-end unified pipeline for accelerating diffusion mod
## NEWS
- ```2025/11/19```: Release [CausalWan2.2 I2V A14B Preview](https://huggingface.co/FastVideo/CausalWan2.2-I2V-A14B-Preview-Diffusers) models, [Blog](https://hao-ai-lab.github.io/blogs/fastvideo_causalwan_preview/) and [Inference Code!](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_self_forcing_causal_wan2_2_i2v.py)
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd.html) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
- ```2025/08/04```: Release [FastWan](https://hao-ai-lab.github.io/FastVideo/distillation/dmd) models and [Sparse-Distillation](https://hao-ai-lab.github.io/blogs/fastvideo_post_training/).
- ```2025/06/14```: Release finetuning and inference code for [VSA](https://arxiv.org/pdf/2505.13389)
- ```2025/04/24```: [FastVideo V1](https://hao-ai-lab.github.io/blogs/fastvideo/) is released!
- ```2025/02/18```: Release the inference code for [Sliding Tile Attention](https://hao-ai-lab.github.io/blogs/sta/).
@@ -125,7 +125,7 @@ For a more detailed guide, please see our [inference quick start](https://hao-ai
## 🤝 Contributing
We welcome all contributions. Please check out our guide [here](https://hao-ai-lab.github.io/FastVideo/contributing/overview/).
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/468).
See details in [development roadmap](https://github.com/hao-ai-lab/FastVideo/issues/899).
## Acknowledgement
We learned and reused code from the following projects:
- [Wan-Video](https://github.com/Wan-Video)
+99 -10
View File
@@ -34,16 +34,16 @@ def pytorch_test(Q, K, V, block_sparse_mask, dO):
)
def block_sparse_kernel_test(Q, K, V, block_sparse_mask, variable_block_sizes, non_pad_index, dO):
def block_sparse_kernel_test(Q, K, V, block_sparse_mask, variable_block_sizes, q_non_pad_index, kv_non_pad_index, q_num_blocks, kv_num_blocks, dO):
Q = Q.detach().requires_grad_()
K = K.detach().requires_grad_()
V = V.detach().requires_grad_()
q_padded = vsa_pad(Q, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
k_padded = vsa_pad(K, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
v_padded = vsa_pad(V, non_pad_index, variable_block_sizes.shape[0], BLOCK_M)
q_padded = vsa_pad(Q, q_non_pad_index, q_num_blocks, BLOCK_M)
k_padded = vsa_pad(K, kv_non_pad_index, kv_num_blocks, BLOCK_M)
v_padded = vsa_pad(V, kv_non_pad_index, kv_num_blocks, BLOCK_M)
output, _= block_sparse_attn(q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes)
output = output[:, :, non_pad_index, :]
output = output[:, :, q_non_pad_index, :]
output.backward(dO)
return output, Q.grad, K.grad, V.grad
@@ -64,7 +64,7 @@ def generate_tensor(shape, dtype, device):
tensor = torch.randn(shape, dtype=dtype, device=device)
return tensor
def generate_variable_block_sizes(num_blocks, min_size=32, max_size=64, device="cuda"):
def generate_variable_block_sizes(num_blocks, min_size=16, max_size=64, device="cuda"):
return torch.randint(min_size, max_size + 1, (num_blocks,), device=device, dtype=torch.int32)
@@ -86,19 +86,21 @@ def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all')
S = int(variable_block_sizes.sum().item())
padded_S = num_blocks * BLOCK_M
non_pad_index = get_non_pad_index(variable_block_sizes, num_blocks, BLOCK_M)
block_mask = generate_block_sparse_mask_for_function(h, num_blocks, k, device)
full_mask = create_full_mask_from_block_mask(block_mask, variable_block_sizes, device)
block_mask = generate_block_sparse_mask_for_function(h, num_blocks, num_blocks, k, device)
full_mask = create_full_mask_from_block_mask(block_mask, variable_block_sizes, variable_block_sizes, device)
for _ in range(num_iterations):
Q = generate_tensor((1, h, S, d), torch.bfloat16, device)
K = generate_tensor((1, h, S, d), torch.bfloat16, device)
V = generate_tensor((1, h, S, d), torch.bfloat16, device)
dO = generate_tensor((1, h, S, d), torch.bfloat16, device)
# print(Q.shape, K.shape, V.shape, dO.shape)
# dO_padded = torch.zeros_like(dO_padded)
# dO_padded[:, :, non_pad_index, :] = dO
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, full_mask, dO)
bs_o, bs_qg, bs_kg, bs_vg = block_sparse_kernel_test(Q, K, V, block_mask.unsqueeze(0), variable_block_sizes,non_pad_index, dO)
bs_o, bs_qg, bs_kg, bs_vg = block_sparse_kernel_test(Q, K, V, block_mask.unsqueeze(0), variable_block_sizes, non_pad_index, non_pad_index, num_blocks, num_blocks, dO)
for name, (pt, bs) in zip(['gQ', 'gK', 'gV', 'gO'], [(pt_qg, bs_qg), (pt_kg, bs_kg), (pt_vg, bs_vg), (pt_o, bs_o)]):
if bs is not None:
diff = pt - bs
@@ -118,6 +120,60 @@ def check_correctness(h, d, num_blocks, k, num_iterations=20, error_mode='all')
return results
def check_correctness_qkdiff(h, d, num_q_blocks, num_kv_blocks, k, num_iterations=20, error_mode='all'):
results = {
'gO': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
'gQ': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
'gK': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
'gV': {'sum_diff': 0.0, 'sum_abs': 0.0, 'max_diff': 0.0},
}
device = "cuda" if torch.cuda.is_available() else "cpu"
q_variable_block_sizes = generate_variable_block_sizes(num_q_blocks, device=device)
kv_variable_block_sizes = generate_variable_block_sizes(num_kv_blocks, device=device)
S_q = int(q_variable_block_sizes.sum().item())
S_kv = int(kv_variable_block_sizes.sum().item())
q_non_pad_index = get_non_pad_index(q_variable_block_sizes, num_q_blocks, BLOCK_M)
kv_non_pad_index = get_non_pad_index(kv_variable_block_sizes, num_kv_blocks, BLOCK_M)
block_mask = generate_block_sparse_mask_for_function(h, num_q_blocks, num_kv_blocks, k, device)
full_mask = create_full_mask_from_block_mask(block_mask, q_variable_block_sizes, kv_variable_block_sizes, device)
for _ in range(num_iterations):
Q = generate_tensor((1, h, S_q, d), torch.bfloat16, device)
K = generate_tensor((1, h, S_kv, d), torch.bfloat16, device)
V = generate_tensor((1, h, S_kv, d), torch.bfloat16, device)
dO = generate_tensor((1, h, S_q, d), torch.bfloat16, device)
# print(Q.shape, K.shape, V.shape, dO.shape)
pt_o, pt_qg, pt_kg, pt_vg = pytorch_test(Q, K, V, full_mask, dO)
bs_o, bs_qg, bs_kg, bs_vg = block_sparse_kernel_test(Q, K, V, block_mask.unsqueeze(0), kv_variable_block_sizes, q_non_pad_index, kv_non_pad_index, num_q_blocks, num_kv_blocks, dO)
for name, (pt, bs) in zip(['gQ', 'gK', 'gV', 'gO'], [(pt_qg, bs_qg), (pt_kg, bs_kg), (pt_vg, bs_vg), (pt_o, bs_o)]):
if bs is not None:
diff = pt - bs
abs_diff = torch.abs(diff)
results[name]['sum_diff'] += torch.sum(abs_diff).item()
results[name]['sum_abs'] += torch.sum(torch.abs(pt)).item()
rel_max_diff = torch.max(abs_diff) / torch.mean(torch.abs(pt))
results[name]['max_diff'] = max(results[name]['max_diff'], rel_max_diff.item())
if torch.cuda.is_available():
torch.cuda.empty_cache()
total_elements_q = h * S_q * d * num_iterations
total_elements_kv = h * S_kv * d * num_iterations
for name, data in results.items():
total_elements = total_elements_q if name in ['gQ', 'gO'] else total_elements_kv
avg_diff = data['sum_diff'] / total_elements
max_diff = data['max_diff']
results[name] = {'avg_diff': avg_diff, 'max_diff': max_diff}
return results
def generate_error_graphs(h, d, error_mode='all'):
test_configs = [
{"num_blocks": 16, "k": 2, "description": "Small sequence"},
@@ -147,10 +203,43 @@ def generate_error_graphs(h, d, error_mode='all'):
print("-" * 150)
def generate_error_graphs_qkdiff(h, d, error_mode='all'):
test_configs = [
{"num_q_blocks": 16, "num_kv_blocks": 32, "k": 2, "description": "Small Q, Med KV"},
{"num_q_blocks": 32, "num_kv_blocks": 16, "k": 4, "description": "Med Q, Small KV"},
{"num_q_blocks": 53, "num_kv_blocks": 32, "k": 6, "description": "Large Q, Med KV"},
{"num_q_blocks": 16, "num_kv_blocks": 48, "k": 2, "description": "Small Q, Large KV"},
{"num_q_blocks": 48, "num_kv_blocks": 16, "k": 2, "description": "Large Q, Small KV"},
]
print(f"\nError Analysis (QK Diff) for h={h}, d={d}, mode={error_mode}")
print("=" * 150)
print(f"{'Config':<20} {'Q Blks':<8} {'KV Blks':<8} {'K':<4} "
f"{'gQ Avg':<12} {'Rel gQ Max':<12} "
f"{'gK Avg':<12} {'Rel gK Max':<12} "
f"{'gV Avg':<12} {'Rel gV Max':<12} "
f"{'gO Avg':<12} {'Rel gO Max':<12}")
print("-" * 150)
for config in test_configs:
num_q_blocks = config["num_q_blocks"]
num_kv_blocks = config["num_kv_blocks"]
k = config["k"]
description = config["description"]
results = check_correctness_qkdiff(h, d, num_q_blocks, num_kv_blocks, k, error_mode=error_mode)
print(f"{description:<20} {num_q_blocks:<8} {num_kv_blocks:<8} {k:<4} "
f"{results['gQ']['avg_diff']:<12.6e} {results['gQ']['max_diff']:<12.6e} "
f"{results['gK']['avg_diff']:<12.6e} {results['gK']['max_diff']:<12.6e} "
f"{results['gV']['avg_diff']:<12.6e} {results['gV']['max_diff']:<12.6e} "
f"{results['gO']['avg_diff']:<12.6e} {results['gO']['max_diff']:<12.6e}")
print("-" * 150)
if __name__ == "__main__":
h, d = 16, 128
print("Block Sparse Attention with Variable Block Sizes Analysis")
print("=" * 60)
for mode in ['backward']:
generate_error_graphs(h, d, error_mode=mode)
print("\nAnalysis completed for all modes.")
generate_error_graphs_qkdiff(h, d, error_mode=mode)
print("\nAnalysis completed for all modes.")
+236
View File
@@ -0,0 +1,236 @@
import os
import sys
from typing import Tuple
import torch
# Make sure we can import from the project root (`vsa`, `tests.utils`, etc.)
CURRENT_DIR = os.path.dirname(os.path.abspath(__file__))
PROJECT_ROOT = os.path.dirname(CURRENT_DIR)
if PROJECT_ROOT not in sys.path:
sys.path.append(PROJECT_ROOT)
if CURRENT_DIR not in sys.path:
sys.path.append(CURRENT_DIR)
from tests.utils import (
generate_block_sparse_mask_for_function,
create_full_mask_from_block_mask,
)
from vsa import block_sparse_attn, BLOCK_M
import test_vsa as ref # reuse helper functions from backward test
def pytorch_forward(
Q: torch.Tensor,
K: torch.Tensor,
V: torch.Tensor,
block_sparse_mask: torch.Tensor,
) -> torch.Tensor:
"""
Dense PyTorch reference forward:
- Q: [1, h, S_q, d]
- K,V: [1, h, S_kv, d]
- block_sparse_mask: [h, S_q, S_kv] bool
"""
q = Q.clone().float()
k = K.clone().float()
v = V.clone().float()
attn = torch.matmul(q, k.transpose(-2, -1)) # [1, h, S_q, S_kv]
attn = attn / (q.size(-1) ** 0.5)
attn = attn.masked_fill(~block_sparse_mask.unsqueeze(0), float("-inf"))
attn = torch.nn.functional.softmax(attn, dim=-1)
out = torch.matmul(attn, v) # [1, h, S_q, d]
return out.to(torch.bfloat16)
def block_sparse_forward_test(
Q: torch.Tensor,
K: torch.Tensor,
V: torch.Tensor,
block_sparse_mask: torch.Tensor,
variable_block_sizes: torch.Tensor,
q_non_pad_index: torch.Tensor,
kv_non_pad_index: torch.Tensor,
q_num_blocks: int,
kv_num_blocks: int,
) -> torch.Tensor:
"""
Forward-only wrapper around `block_sparse_attn`, mirroring `block_sparse_kernel_test`
but without any backward / grad logic.
"""
Q = Q.detach()
K = K.detach()
V = V.detach()
q_padded = ref.vsa_pad(Q, q_non_pad_index, q_num_blocks, BLOCK_M)
k_padded = ref.vsa_pad(K, kv_non_pad_index, kv_num_blocks, BLOCK_M)
v_padded = ref.vsa_pad(V, kv_non_pad_index, kv_num_blocks, BLOCK_M)
out_padded, _ = block_sparse_attn(
q_padded, k_padded, v_padded, block_sparse_mask, variable_block_sizes
)
# Remove padding on the query side
out = out_padded[:, :, q_non_pad_index, :]
return out
def run_forward_equal_qk(
h: int = 16,
d: int = 128,
num_blocks: int = 16,
k: int = 2,
num_iterations: int = 5,
) -> Tuple[float, float]:
"""
Forward-only correctness test for the case S_q == S_kv.
Mirrors `check_correctness` but only compares forward outputs.
"""
assert torch.cuda.is_available(), "VSA kernels require CUDA"
device = "cuda"
variable_block_sizes = ref.generate_variable_block_sizes(
num_blocks, device=device
)
S = int(variable_block_sizes.sum().item())
non_pad_index = ref.get_non_pad_index(
variable_block_sizes, num_blocks, BLOCK_M
)
block_mask = generate_block_sparse_mask_for_function(
h, num_blocks, num_blocks, k, device
)
full_mask = create_full_mask_from_block_mask(
block_mask, variable_block_sizes, variable_block_sizes, device
)
print(f"[qkequal] h: {h}, d: {d}, num_blocks: {num_blocks}, k: {k}")
print(f"[qkequal] variable_block_sizes: {variable_block_sizes}, non_pad_index: {non_pad_index.shape}, block_mask: {block_mask.shape}, full_mask: {full_mask.shape}")
sum_diff = 0.0
sum_abs = 0.0
max_rel_diff = 0.0
for i in range(num_iterations):
Q = ref.generate_tensor((1, h, S, d), torch.bfloat16, device)
K = ref.generate_tensor((1, h, S, d), torch.bfloat16, device)
V = ref.generate_tensor((1, h, S, d), torch.bfloat16, device)
if i == 0: print(f"[qkequal] Q: {Q.shape}, K: {K.shape}, V: {V.shape}, full_mask: {full_mask.shape}")
if i == 0: print(f"[qkequal] block_mask: {block_mask.shape}")
pt_o = pytorch_forward(Q, K, V, full_mask)
bs_o = block_sparse_forward_test(
Q,
K,
V,
block_mask.unsqueeze(0),
variable_block_sizes,
non_pad_index,
non_pad_index,
num_blocks,
num_blocks,
)
diff = (pt_o - bs_o).abs()
sum_diff += diff.sum().item()
sum_abs += pt_o.abs().sum().item()
rel_max = diff.max() / (pt_o.abs().mean() + 1e-6)
max_rel_diff = max(max_rel_diff, rel_max.item())
total_elems = h * S * d * num_iterations
avg_abs_err = sum_diff / total_elems
return avg_abs_err, max_rel_diff
def run_forward_qk_diff(
h: int = 16,
d: int = 128,
num_q_blocks: int = 16,
num_kv_blocks: int = 32,
k: int = 2,
num_iterations: int = 5,
) -> Tuple[float, float]:
"""
Forward-only correctness test for the case S_q != S_kv.
NOTE:
- The Triton backend supports different Q/KV logical lengths via padding.
- The SM90 (H100) CUDA backend currently assumes the same number of blocks
for Q and KV, so we skip this test there.
"""
assert torch.cuda.is_available(), "VSA kernels require CUDA"
device = "cuda"
q_variable_block_sizes = ref.generate_variable_block_sizes(
num_q_blocks, device=device
)
kv_variable_block_sizes = ref.generate_variable_block_sizes(
num_kv_blocks, device=device
)
S_q = int(q_variable_block_sizes.sum().item())
S_kv = int(kv_variable_block_sizes.sum().item())
q_non_pad_index = ref.get_non_pad_index(
q_variable_block_sizes, num_q_blocks, BLOCK_M
)
kv_non_pad_index = ref.get_non_pad_index(
kv_variable_block_sizes, num_kv_blocks, BLOCK_M
)
block_mask = generate_block_sparse_mask_for_function(
h, num_q_blocks, num_kv_blocks, k, device
)
full_mask = create_full_mask_from_block_mask(
block_mask, q_variable_block_sizes, kv_variable_block_sizes, device
)
sum_diff = 0.0
sum_abs = 0.0
max_rel_diff = 0.0
for _ in range(num_iterations):
Q = ref.generate_tensor((1, h, S_q, d), torch.bfloat16, device)
K = ref.generate_tensor((1, h, S_kv, d), torch.bfloat16, device)
V = ref.generate_tensor((1, h, S_kv, d), torch.bfloat16, device)
pt_o = pytorch_forward(Q, K, V, full_mask)
bs_o = block_sparse_forward_test(
Q,
K,
V,
block_mask.unsqueeze(0),
kv_variable_block_sizes,
q_non_pad_index,
kv_non_pad_index,
num_q_blocks,
num_kv_blocks,
)
diff = (pt_o - bs_o).abs()
sum_diff += diff.sum().item()
sum_abs += pt_o.abs().sum().item()
rel_max = diff.max() / (pt_o.abs().mean() + 1e-6)
max_rel_diff = max(max_rel_diff, rel_max.item())
total_elems = h * S_q * d * num_iterations
avg_abs_err = sum_diff / total_elems
return avg_abs_err, max_rel_diff
if __name__ == "__main__":
h, d = 16, 128
print("Forward Block Sparse Attention Check (QK Equal)")
print("=" * 80)
avg_err_eq, max_rel_eq = run_forward_equal_qk(h, d, num_blocks=32, k=2)
print(f"QK equal: avg |ΔO| = {avg_err_eq:.6e}, max rel ΔO = {max_rel_eq:.6e}")
print("\nForward Block Sparse Attention Check (QK Different)")
print("=" * 80)
avg_err_diff, max_rel_diff = run_forward_qk_diff(
h, d, num_q_blocks=32, num_kv_blocks=48, k=2
)
print(
f"QK diff: avg |ΔO| = {avg_err_diff:.6e}, max rel ΔO = {max_rel_diff:.6e}"
)
+27 -21
View File
@@ -1,54 +1,60 @@
import torch
def generate_block_sparse_mask_for_function(h, num_blocks, k, device="cuda"):
def generate_block_sparse_mask_for_function(h, num_q_blocks, num_kv_blocks, k, device="cuda"):
"""
Generate block sparse mask of shape [h, num_blocks, num_blocks].
Generate block sparse mask of shape [h, num_q_blocks, num_kv_blocks].
Args:
h: number of heads
num_blocks: number of blocks
num_q_blocks: number of query blocks
num_kv_blocks: number of key/value blocks
k: number of kv blocks each q block attends to
device: device to create tensors on
Returns:
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
block_sparse_mask: [h, num_q_blocks, num_kv_blocks] bool tensor
"""
k = min(k, num_blocks)
scores = torch.rand(h, num_blocks, num_blocks, device=device)
k = min(k, num_kv_blocks)
scores = torch.rand(h, num_q_blocks, num_kv_blocks, device=device)
_, indices = torch.topk(scores, k, dim=-1)
block_sparse_mask = torch.zeros(h, num_blocks, num_blocks, dtype=torch.bool, device=device)
block_sparse_mask = torch.zeros(h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device=device)
block_sparse_mask = block_sparse_mask.scatter_(2, indices, 1).bool()
return block_sparse_mask
def create_full_mask_from_block_mask(block_sparse_mask, variable_block_sizes, device="cuda"):
def create_full_mask_from_block_mask(block_sparse_mask, q_variable_block_sizes,
kv_variable_block_sizes, device="cuda"):
"""
Convert block-level sparse mask to full attention mask.
Args:
block_sparse_mask: [h, num_blocks, num_blocks] bool tensor
variable_block_sizes: [num_blocks] tensor
block_sparse_mask: [h, num_q_blocks, num_kv_blocks] bool tensor
q_variable_block_sizes: [num_q_blocks] tensor
kv_variable_block_sizes: [num_kv_blocks] tensor
device: device to create tensors on
Returns:
full_mask: [h, S, S] bool tensor where S = total sequence length
full_mask: [h, S_q, S_kv] bool tensor where S = total sequence length
"""
h, num_blocks, _ = block_sparse_mask.shape
total_seq_len = variable_block_sizes.sum().item()
cumsum = torch.cat([torch.tensor([0], device=device), variable_block_sizes.cumsum(dim=0)[:-1]])
h, num_q_blocks, num_kv_blocks = block_sparse_mask.shape
total_q_seq_len = q_variable_block_sizes.sum().item()
total_kv_seq_len = kv_variable_block_sizes.sum().item()
q_cumsum = torch.cat([torch.tensor([0], device=device), q_variable_block_sizes.cumsum(dim=0)[:-1]])
kv_cumsum = torch.cat([torch.tensor([0], device=device), kv_variable_block_sizes.cumsum(dim=0)[:-1]])
full_mask = torch.zeros(h, total_seq_len, total_seq_len, dtype=torch.bool, device=device)
full_mask = torch.zeros(h, total_q_seq_len, total_kv_seq_len, dtype=torch.bool, device=device)
for head in range(h):
for q_block in range(num_blocks):
q_start = cumsum[q_block]
q_end = q_start + variable_block_sizes[q_block]
for q_block in range(num_q_blocks):
q_start = q_cumsum[q_block]
q_end = q_start + q_variable_block_sizes[q_block]
for kv_block in range(num_blocks):
for kv_block in range(num_kv_blocks):
if block_sparse_mask[head, q_block, kv_block]:
kv_start = cumsum[kv_block]
kv_end = kv_start + variable_block_sizes[kv_block]
kv_start = kv_cumsum[kv_block]
kv_end = kv_start + kv_variable_block_sizes[kv_block]
full_mask[head, q_start:q_end, kv_start:kv_end] = True
return full_mask
@@ -672,23 +672,32 @@ block_sparse_attention_forward(
torch::Tensor v,
torch::Tensor q2k_block_sparse_index,
torch::Tensor q2k_block_sparse_num,
torch::Tensor block_size
torch::Tensor kv_block_size
)
{
CHECK_INPUT(q);
CHECK_INPUT(k);
CHECK_INPUT(v);
// q shape: (batch, qo_heads, q_seq_len, head_dim)
// k shape: (batch, kv_heads, kv_seq_len, head_dim)
// v shape: (batch, kv_heads, kv_seq_len, head_dim)
// q2k_block_sparse_index shape: (batch, qo_heads, num_q_blocks, max_kv_blocks_per_q)
// q2k_block_sparse_num shape: (batch, qo_heads, num_q_blocks)
// kv_block_size shape: (num_kv_blocks) This does not need other dimensions because across all batch/heads the padding is the same.
auto batch = q.size(0);
auto seq_len = q.size(2);
auto q_seq_len = q.size(2);
auto kv_seq_len = k.size(2);
auto head_dim = q.size(3);
auto qo_heads = q.size(1);
auto kv_heads = k.size(1);
auto max_kv_blocks_per_q = q2k_block_sparse_index.size(3);
auto num_q_blocks = block_size.size(0);
auto num_q_blocks = q2k_block_sparse_index.size(2);
auto num_kv_blocks = kv_block_size.size(0);
TORCH_CHECK(batch==1, "Batch size dim will be removed in the future, please set batch to 1");
TORCH_CHECK(num_q_blocks * 64 == seq_len, "This kernel supports variable block size, but it assumes the input sequence is properly padded.");
TORCH_CHECK(num_q_blocks == q2k_block_sparse_index.size(2), "Number of Q blocks does not match between q2k_block_sparse_index and block_size");
TORCH_CHECK(num_q_blocks * BLOCK_M == q_seq_len, "This kernel supports variable q block size, but it assumes the input sequence is properly padded.");
TORCH_CHECK(num_kv_blocks * BLOCK_M == kv_seq_len, "This kernel supports variable kv block size, but it assumes the input sequence is properly padded.");
// check to see that these dimensions match for all inputs
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
@@ -696,11 +705,8 @@ block_sparse_attention_forward(
TORCH_CHECK(q2k_block_sparse_index.size(0) == batch, "q2k_block_sparse_index batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(q2k_block_sparse_num.size(0) == batch, "q2k_block_sparse_num batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(q2k_block_sparse_index.size(2) == seq_len / BLOCK_M, "q2k_block_sparse_index idx 2 - must match seq_len / BLOCK_M");
TORCH_CHECK(q2k_block_sparse_num.size(2) == seq_len / BLOCK_M, "q2k_block_sparse_num idx 2 - must match seq_len / BLOCK_M");
TORCH_CHECK(v.size(2) == kv_seq_len, "V sequence length dimension - idx 2 - must match K inputs");
TORCH_CHECK(q2k_block_sparse_num.size(2) == num_q_blocks, "q2k_block_sparse_num idx 2 - must match num_q_blocks");
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
@@ -727,12 +733,12 @@ block_sparse_attention_forward(
// for the returned outputs
torch::Tensor o = torch::empty({static_cast<const uint>(batch),
static_cast<const uint>(qo_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(q_seq_len),
static_cast<const uint>(head_dim)}, v.options());
torch::Tensor l_vec = torch::empty({static_cast<const uint>(batch),
static_cast<const uint>(qo_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(q_seq_len),
static_cast<const uint>(1)},
torch::TensorOptions().dtype(torch::kFloat).device(q.device()).memory_format(at::MemoryFormat::Contiguous));
@@ -762,11 +768,11 @@ block_sparse_attention_forward(
using globals = fwd_globals<64>;
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
globals g{
qg_arg,
@@ -774,17 +780,17 @@ block_sparse_attention_forward(
vg_arg,
lg_arg,
og_arg,
static_cast<int>(seq_len),
static_cast<int>(q_seq_len),
static_cast<int>(hr),
static_cast<int>(max_kv_blocks_per_q),
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
reinterpret_cast<int32_t*>(block_size.data_ptr())
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())
};
constexpr int mem_size = 54000;
dim3 grid(seq_len/(64), qo_heads, batch);
dim3 grid(q_seq_len/(BLOCK_M), qo_heads, batch);
cudaFuncSetAttribute(
fwd_attend_ker<64>,
@@ -813,11 +819,11 @@ block_sparse_attention_forward(
using globals = fwd_globals<128>;
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
q_global qg_arg{d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
k_global kg_arg{d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
v_global vg_arg{d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
l_global lg_arg{d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
o_global og_arg{d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
globals g{
qg_arg,
@@ -825,17 +831,17 @@ block_sparse_attention_forward(
vg_arg,
lg_arg,
og_arg,
static_cast<int>(seq_len),
static_cast<int>(q_seq_len),
static_cast<int>(hr),
static_cast<int>(max_kv_blocks_per_q),
reinterpret_cast<int32_t*>(q2k_block_sparse_index.data_ptr()),
reinterpret_cast<int32_t*>(q2k_block_sparse_num.data_ptr()),
reinterpret_cast<int32_t*>(block_size.data_ptr())
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())
};
constexpr int mem_size = 54000;
dim3 grid(seq_len/(64), qo_heads, batch);
dim3 grid(q_seq_len/(BLOCK_M), qo_heads, batch);
cudaFuncSetAttribute(
fwd_attend_ker<128>,
@@ -862,7 +868,7 @@ block_sparse_attention_backward(torch::Tensor q,
torch::Tensor og,
torch::Tensor k2q_block_sparse_index,
torch::Tensor k2q_block_sparse_num,
torch::Tensor block_size)
torch::Tensor kv_block_size)
{
CHECK_INPUT(q);
CHECK_INPUT(k);
@@ -871,11 +877,23 @@ block_sparse_attention_backward(torch::Tensor q,
CHECK_INPUT(o);
CHECK_INPUT(og);
// q: [batch, qo_heads, q_seq_len, head_dim]
// k: [batch, kv_heads, kv_seq_len, head_dim]
// v: [batch, kv_heads, kv_seq_len, head_dim]
// o: [batch, qo_heads, q_seq_len, head_dim]
// l_vec: [batch, qo_heads, q_seq_len, 1]
// og: [batch, qo_heads, q_seq_len, head_dim]
// k2q_block_sparse_index: [batch, kv_heads, num_kv_blocks, max_num_q_blocks]
// k2q_block_sparse_num: [batch, kv_heads, num_kv_blocks]
// kv_block_size: [num_kv_blocks]
auto batch = q.size(0);
auto seq_len = q.size(2);
auto q_seq_len = q.size(2);
auto kv_seq_len = k.size(2);
auto head_dim = q.size(3);
auto max_q_blocks_per_kv = k2q_block_sparse_index.size(3);
TORCH_CHECK(k2q_block_sparse_index.size(2) == block_size.size(0), "k2q_block_sparse_index.size(2) must match block_size.size(0)");
auto num_kv_blocks = kv_block_size.size(0);
TORCH_CHECK(k2q_block_sparse_index.size(2) == num_kv_blocks, "k2q_block_sparse_index.size(2) must match num_kv_blocks (kv_block_size.size(0))");
// check to see that these dimensions match for all inputs
TORCH_CHECK(q.size(0) == batch, "Q batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(k.size(0) == batch, "K batch dimension - idx 0 - must match for all inputs");
@@ -886,23 +904,18 @@ block_sparse_attention_backward(torch::Tensor q,
TORCH_CHECK(k2q_block_sparse_index.size(0) == batch, "k2q_block_sparse_index batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(k2q_block_sparse_num.size(0) == batch, "k2q_block_sparse_num batch dimension - idx 0 - must match for all inputs");
TORCH_CHECK(q.size(2) == seq_len, "Q sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(k.size(2) == seq_len, "K sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(v.size(2) == seq_len, "V sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(l_vec.size(2) == seq_len, "L sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(o.size(2) == seq_len, "O sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(og.size(2) == seq_len, "OG sequence length dimension - idx 2 - must match for all inputs");
TORCH_CHECK(k2q_block_sparse_index.size(2) == seq_len / BLOCK_N, "k2q_block_sparse_index idx 2 - must match seq_len / BLOCK_N");
TORCH_CHECK(k2q_block_sparse_num.size(2) == seq_len / BLOCK_N, "k2q_block_sparse_num idx 2 - must match seq_len / BLOCK_N");
TORCH_CHECK(v.size(2) == kv_seq_len, "V sequence length dimension - idx 2 - must match K sequence length");
TORCH_CHECK(l_vec.size(2) == q_seq_len, "L sequence length dimension - idx 2 - must match Q sequence length");
TORCH_CHECK(o.size(2) == q_seq_len, "O sequence length dimension - idx 2 - must match Q sequence length");
TORCH_CHECK(og.size(2) == q_seq_len, "OG sequence length dimension - idx 2 - must match Q sequence length");
TORCH_CHECK(k2q_block_sparse_index.size(2) == num_kv_blocks, "k2q_block_sparse_index idx 2 - must match num_kv_blocks (kv_block_size.size(0))");
TORCH_CHECK(k2q_block_sparse_num.size(2) == num_kv_blocks, "k2q_block_sparse_num idx 2 - must match num_kv_blocks (kv_block_size.size(0))");
TORCH_CHECK(q.size(3) == head_dim, "Q head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(k.size(3) == head_dim, "K head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(v.size(3) == head_dim, "V head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(o.size(3) == head_dim, "O head dimension - idx 3 - must match for all non-vector inputs");
TORCH_CHECK(og.size(3) == head_dim, "OG head dimension - idx 3 - must match for all non-vector inputs");
auto qo_heads = q.size(1);
auto kv_heads = k.size(1);
@@ -929,20 +942,20 @@ block_sparse_attention_backward(torch::Tensor q,
torch::Tensor qg = torch::zeros({static_cast<const uint>(batch),
static_cast<const uint>(qo_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(q_seq_len),
static_cast<const uint>(head_dim)}, l_vec.options());
torch::Tensor kg = torch::zeros({static_cast<const uint>(batch),
static_cast<const uint>(kv_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(kv_seq_len),
static_cast<const uint>(head_dim)}, l_vec.options());
torch::Tensor vg = torch::zeros({static_cast<const uint>(batch),
static_cast<const uint>(kv_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(kv_seq_len),
static_cast<const uint>(head_dim)}, l_vec.options());
torch::Tensor d_vec = torch::empty({static_cast<const uint>(batch),
static_cast<const uint>(qo_heads),
static_cast<const uint>(seq_len),
static_cast<const uint>(q_seq_len),
static_cast<const uint>(1)}, l_vec.options());
float* qg_ptr = qg.data_ptr<float>();
@@ -971,7 +984,7 @@ block_sparse_attention_backward(torch::Tensor q,
// cudaStreamSynchronize(stream);
// TORCH_CHECK(seq_len % (4*kittens::TILE_DIM*4) == 0, "sequence length must be divisible by 256");
dim3 grid_bwd(seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
dim3 grid_bwd(q_seq_len/(PREP_NUM_WARPS*kittens::TILE_ROW_DIM<bf16>*4), qo_heads, batch);
if (head_dim == 64) {
using og_tile = st_bf<4*16, 64>;
@@ -984,9 +997,9 @@ block_sparse_attention_backward(torch::Tensor q,
using bwd_prep_globals = bwd_prep_globals<64>;
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
bwd_prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
@@ -1023,15 +1036,15 @@ block_sparse_attention_backward(torch::Tensor q,
using bwd_global_args = bwd_globals<64>;
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 64U};
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 64U};
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 64U};
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 64U};
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
bwd_global_args bwd_global{bwd_q_arg,
bwd_k_arg,
@@ -1042,14 +1055,14 @@ block_sparse_attention_backward(torch::Tensor q,
bwd_vg_arg,
bwd_l_arg,
bwd_d_arg,
static_cast<int>(seq_len),
static_cast<int>(kv_seq_len), // N is not used in the kernel
static_cast<int>(hr),
static_cast<int>(max_q_blocks_per_kv),
reinterpret_cast<int32_t*>(k2q_block_sparse_index.data_ptr()),
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr()),
reinterpret_cast<int32_t*>(block_size.data_ptr())};
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())};
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
dim3 grid_bwd_2(kv_seq_len/BLOCK_N, qo_heads, batch);
threads = 128;
//cudadevicesynchronize();
@@ -1088,9 +1101,9 @@ block_sparse_attention_backward(torch::Tensor q,
using bwd_prep_globals = bwd_prep_globals<128>;
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
og_global prep_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
o_global prep_o_arg {d_o, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
d_global prep_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
bwd_prep_globals bwd_g{prep_og_arg, prep_o_arg, prep_d_arg};
@@ -1127,15 +1140,15 @@ block_sparse_attention_backward(torch::Tensor q,
using bwd_global_args = bwd_globals<128>;
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(seq_len), 128U};
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(seq_len), 128U};
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(seq_len)};
bwd_q_global bwd_q_arg {d_q, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
bwd_k_global bwd_k_arg {d_k, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
bwd_v_global bwd_v_arg {d_v, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
bwd_og_global bwd_og_arg{d_og, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
bwd_qg_global bwd_qg_arg{d_qg, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), static_cast<unsigned int>(q_seq_len), 128U};
bwd_kg_global bwd_kg_arg{d_kg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
bwd_vg_global bwd_vg_arg{d_vg, static_cast<unsigned int>(batch), static_cast<unsigned int>(kv_heads), static_cast<unsigned int>(kv_seq_len), 128U};
bwd_l_global bwd_l_arg {d_l, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
bwd_d_global bwd_d_arg {d_d, static_cast<unsigned int>(batch), static_cast<unsigned int>(qo_heads), 1U, static_cast<unsigned int>(q_seq_len)};
bwd_global_args bwd_global{bwd_q_arg,
bwd_k_arg,
@@ -1146,14 +1159,14 @@ block_sparse_attention_backward(torch::Tensor q,
bwd_vg_arg,
bwd_l_arg,
bwd_d_arg,
static_cast<int>(seq_len),
static_cast<int>(kv_seq_len), // N is not used in the kernel
static_cast<int>(hr),
static_cast<int>(max_q_blocks_per_kv),
reinterpret_cast<int32_t*>(k2q_block_sparse_index.data_ptr()),
reinterpret_cast<int32_t*>(k2q_block_sparse_num.data_ptr()),
reinterpret_cast<int32_t*>(block_size.data_ptr())};
reinterpret_cast<int32_t*>(kv_block_size.data_ptr())};
dim3 grid_bwd_2(seq_len/64, qo_heads, batch);
dim3 grid_bwd_2(kv_seq_len/BLOCK_N, qo_heads, batch);
threads = 128;
//cudadevicesynchronize();
+2 -2
View File
@@ -1,4 +1,4 @@
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
ENV DEBIAN_FRONTEND=noninteractive
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp310-cp310-linux_x86_64.whl
COPY . .
+2 -2
View File
@@ -1,4 +1,4 @@
FROM nvidia/cuda:12.8.0-devel-ubuntu22.04
FROM nvidia/cuda:12.8.0-cudnn-devel-ubuntu22.04
ENV DEBIAN_FRONTEND=noninteractive
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp311-cp311-linux_x86_64.whl
COPY . .
+1 -1
View File
@@ -43,7 +43,7 @@ RUN source $HOME/.local/bin/env && \
source /opt/venv/bin/activate && \
uv pip install --no-cache-dir --upgrade pip && \
uv pip install --no-cache-dir .[dev] && \
uv pip install --no-cache-dir flash-attn==2.8.3 --no-build-isolation
uv pip install --no-cache-dir https://github.com/mjun0812/flash-attention-prebuild-wheels/releases/download/v0.5.4/flash_attn-2.8.3%2Bcu128torch2.9-cp312-cp312-linux_x86_64.whl
COPY . .
+1 -1
View File
@@ -13,7 +13,7 @@ We provide two distilled models:
Both models are trained on **61×448×832** resolution but support generating videos with **any resolution** (1.3B model mainly support 480P, 14B model support 480P and 720P, quality may degrade for different resolutions).
## ⚙️ Inference
First install [VSA](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation.html). Set `MODEL_BASE` to your own model path and run:
First install [VSA](https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation). Set `MODEL_BASE` to your own model path and run:
```bash
bash scripts/inference/v1_inference_wan_dmd.sh
+1 -1
View File
@@ -6,7 +6,7 @@ The `VideoGenerator` class provides the primary Python interface for doing offli
- Python 3.10-3.12
## Installation
If you have not installed FastVideo, please following these [instructions](https://hao-ai-lab.github.io/FastVideo/getting_started/installation.html) first.
If you have not installed FastVideo, please following these [instructions](https://hao-ai-lab.github.io/FastVideo/getting_started/installation) first.
## Usage
The first script in this example shows the most basic usage of FastVideo. If you are new to Python and FastVideo, you should start here.
+42
View File
@@ -0,0 +1,42 @@
from fastvideo import VideoGenerator
import json
# from fastvideo.configs.sample import SamplingParam
OUTPUT_PATH = "video_samples_hy15"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
generator = VideoGenerator.from_pretrained(
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v",
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True, # set to false if low CPU RAM or hit obscure "CUDA error: Invalid argument"
# image_encoder_cpu_offload=False,
)
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, negative_prompt="", num_frames=81, fps=16)
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, negative_prompt="", num_frames=81, fps=16)
if __name__ == "__main__":
main()
@@ -0,0 +1,75 @@
from fastvideo import VideoGenerator
from fastvideo.configs.pipelines.wan import MatrixGameI2V480PConfig
from fastvideo.models.dits.matrix_game.utils import create_action_presets
import torch
# Available variants: "base_distilled_model", "gta_distilled_model", "templerun_distilled_model"
# Each variant has different keyboard_dim:
# - base_distilled_model: keyboard_dim=4
# - gta_distilled_model: keyboard_dim=2
# - templerun_distilled_model: keyboard_dim=7 (keyboard only, no mouse)
MODEL_VARIANT = "base_distilled_model"
# Variant-specific settings
VARIANT_CONFIG = {
"base_distilled_model": {
"model_path": "FastVideo/Matrix-Game-2.0-Base-Diffusers",
"keyboard_dim": 4,
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0000.png",
},
"gta_distilled_model": {
"model_path": "FastVideo/Matrix-Game-2.0-GTA-Diffusers",
"keyboard_dim": 2,
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/gta_drive/0000.png",
},
"templerun_distilled_model": {
"model_path": "FastVideo/Matrix-Game-2.0-TempleRun-Diffusers",
"keyboard_dim": 7,
"image_url": "https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/temple_run/0000.png",
},
}
OUTPUT_PATH = "video_samples_matrixgame2"
def main():
# FastVideo will automatically use the optimal default arguments for the
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
config = VARIANT_CONFIG[MODEL_VARIANT]
generator = VideoGenerator.from_pretrained(
config["model_path"],
# FastVideo will automatically handle distributed setup
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
pin_cpu_memory=True,
# image_encoder_cpu_offload=False,
)
num_frames = 597
actions = create_action_presets(num_frames, keyboard_dim=config["keyboard_dim"])
grid_sizes = torch.tensor([150, 44, 80])
generator.generate_video(
prompt="",
image_path=config["image_url"],
mouse_cond=actions["mouse"].unsqueeze(0),
keyboard_cond=actions["keyboard"].unsqueeze(0),
grid_sizes=grid_sizes,
num_frames=num_frames,
height=352,
width=640,
num_inference_steps=50,
output_path=OUTPUT_PATH,
save_video=True,
)
if __name__ == "__main__":
main()
@@ -0,0 +1,107 @@
#!/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/combined_parquet_dataset/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=4
# export CUDA_VISIBLE_DEVICES=4,5
mkdir -p ../profiler_traces/wan_t2v_finetune/
# Torch Profiler Configuration
export FASTVIDEO_TORCH_PROFILE_REGIONS="profiler_region_training_train_one_step"
export FASTVIDEO_TORCH_PROFILER_RECORD_SHAPES=1
export FASTVIDEO_TORCH_PROFILER_WITH_STACK=1
export FASTVIDEO_TORCH_PROFILER_WITH_FLOPS=1
export FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY=1
export FASTVIDEO_TORCH_PROFILER_WAIT_STEPS=2
export FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS=1
export FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS=1
export FASTVIDEO_TORCH_PROFILER_DIR="../profiler_traces/wan_t2v_finetune/"
# Training arguments
training_args=(
--tracker_project_name "wan_t2v_finetune"
--output_dir "checkpoints/wan_t2v_finetune"
--max_train_steps 5000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 8
--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 $NUM_GPUS
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
)
# 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 200
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 5e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_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"
# --resume_from_checkpoint "checkpoints/wan_t2v_finetune/checkpoint-2500"
)
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
--master_port 29502 \
fastvideo/training/wan_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
+47 -8
View File
@@ -1,8 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
import torch
import torch.nn.functional as F
from flash_attn import flash_attn_func as flash_attn_2_func
from dataclasses import dataclass
try:
from flash_attn_interface import flash_attn_func as flash_attn_3_func
@@ -46,6 +47,29 @@ class FlashAttentionBackend(AttentionBackend):
raise NotImplementedError
@dataclass
class FlashAttnMetadata(AttentionMetadata):
current_timestep: int
attn_mask: torch.Tensor | None = None
class FlashAttnMetadataBuilder(AttentionMetadataBuilder):
def __init__(self):
pass
def prepare(self):
pass
def build( # type: ignore
self,
current_timestep: int,
attn_mask: torch.Tensor,
) -> FlashAttnMetadata:
return FlashAttnMetadata(current_timestep=current_timestep,
attn_mask=attn_mask)
class FlashAttentionImpl(AttentionImpl):
def __init__(
@@ -66,12 +90,27 @@ class FlashAttentionImpl(AttentionImpl):
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: AttentionMetadata,
attn_metadata: FlashAttnMetadata,
):
output = flash_attn_func(
query, # type: ignore[no-untyped-call]
key,
value,
softmax_scale=self.softmax_scale,
causal=self.causal)
if attn_metadata is not None and hasattr(
attn_metadata,
"attn_mask") and attn_metadata.attn_mask is not None:
from fastvideo.attention.utils.flash_attn_no_pad import flash_attn_no_pad
attn_mask = attn_metadata.attn_mask
qkv = torch.stack([query, key, value], dim=2)
attn_mask = F.pad(attn_mask, (qkv.shape[1] - attn_mask.shape[1], 0),
value=True)
output = flash_attn_no_pad(qkv,
attn_mask,
causal=False,
dropout_p=0,
softmax_scale=None)
else:
output = flash_attn_func(
query, # type: ignore[no-untyped-call]
key,
value,
softmax_scale=self.softmax_scale,
causal=self.causal)
return output
+29 -4
View File
@@ -1,9 +1,10 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from dataclasses import dataclass
from fastvideo.attention.backends.abstract import ( # FlashAttentionMetadata,
AttentionBackend, AttentionImpl, AttentionMetadata)
AttentionBackend, AttentionImpl, AttentionMetadata,
AttentionMetadataBuilder)
from fastvideo.logger import init_logger
logger = init_logger(__name__)
@@ -30,6 +31,29 @@ class SDPABackend(AttentionBackend):
# return FlashAttentionMetadata
@dataclass
class SDPAMetadata(AttentionMetadata):
current_timestep: int
attn_mask: torch.Tensor | None = None
class SDPAMetadataBuilder(AttentionMetadataBuilder):
def __init__(self):
pass
def prepare(self):
pass
def build( # type: ignore
self,
current_timestep: int,
attn_mask: torch.Tensor,
) -> SDPAMetadata:
return SDPAMetadata(current_timestep=current_timestep,
attn_mask=attn_mask)
class SDPAImpl(AttentionImpl):
def __init__(
@@ -51,14 +75,15 @@ class SDPAImpl(AttentionImpl):
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: AttentionMetadata,
attn_metadata: SDPAMetadata,
) -> torch.Tensor:
# transpose to bs, heads, seq_len, head_dim
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
attn_mask = attn_metadata.attn_mask if attn_metadata is not None else None
attn_kwargs = {
"attn_mask": None,
"attn_mask": attn_mask,
"dropout_p": self.dropout,
"is_causal": self.causal,
"scale": self.softmax_scale
+62 -4
View File
@@ -11,6 +11,7 @@ from fastvideo.distributed.parallel_state import (get_sp_parallel_rank,
from fastvideo.forward_context import ForwardContext, get_forward_context
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.utils import get_compute_dtype
from fastvideo.layers.rotary_embedding import _apply_rotary_emb
class DistributedAttention(nn.Module):
@@ -64,6 +65,8 @@ class DistributedAttention(nn.Module):
replicated_q: torch.Tensor | None = None,
replicated_k: torch.Tensor | None = None,
replicated_v: torch.Tensor | None = None,
freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
attention_mask: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""Forward pass for distributed attention.
@@ -74,6 +77,7 @@ class DistributedAttention(nn.Module):
replicated_q (Optional[torch.Tensor]): Replicated query tensor, typically for text tokens
replicated_k (Optional[torch.Tensor]): Replicated key tensor
replicated_v (Optional[torch.Tensor]): Replicated value tensor
attention_mask (Optional[torch.Tensor]): Attention mask [batch_size, seq_len]
Returns:
Tuple[torch.Tensor, Optional[torch.Tensor]]: A tuple containing:
@@ -91,12 +95,30 @@ class DistributedAttention(nn.Module):
ctx_attn_metadata = forward_context.attn_metadata
# Stack QKV
qkv = torch.cat([q, k, v], dim=0) # [3, seq_len, num_heads, head_dim]
qkv = torch.cat([q, k, v],
dim=0) # [3*batch, seq_len, num_heads, head_dim]
# Redistribute heads across sequence dimension
qkv = sequence_model_parallel_all_to_all_4D(qkv,
scatter_dim=2,
gather_dim=1)
# After all-to-all, each rank has the full sequence but only a subset of heads
# The attention mask should now apply to the full sequence length
# Since mask is [batch, full_seq_len], it's already in the correct format
# LOAY TODO, instead of slicing repeatedly maintain an original qkv and rewrite into that
valid_seq_len = None
if attention_mask is not None:
valid_seq_len = (attention_mask[0] == 1).sum().item()
qkv = qkv[:, :valid_seq_len, :, :]
if freqs_cis is not None:
cos, sin = freqs_cis
qkv[:batch_size * 2] = _apply_rotary_emb(qkv[:batch_size * 2],
cos,
sin,
is_neox_style=False)
# Apply backend-specific preprocess_qkv
qkv = self.attn_impl.preprocess_qkv(qkv, ctx_attn_metadata)
@@ -119,17 +141,23 @@ class DistributedAttention(nn.Module):
# Redistribute back if using sequence parallelism
replicated_output = None
if replicated_q is not None:
replicated_output = output[:, seq_len * world_size:]
output = output[:, :seq_len * world_size]
split_idx = seq_len * world_size if valid_seq_len is None else valid_seq_len
replicated_output = output[:, split_idx:]
output = output[:, :split_idx]
# TODO: make this asynchronous
replicated_output = sequence_model_parallel_all_gather(
replicated_output.contiguous(), dim=2)
# Apply backend-specific postprocess_output
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
if attention_mask is not None:
pad_len = (attention_mask[0] == 0).sum().item()
output = torch.nn.functional.pad(output, (0, 0, 0, 0, 0, pad_len))
output = sequence_model_parallel_all_to_all_4D(output,
scatter_dim=1,
gather_dim=2)
return output, replicated_output
@@ -147,6 +175,8 @@ class DistributedAttention_VSA(DistributedAttention):
replicated_k: torch.Tensor | None = None,
replicated_v: torch.Tensor | None = None,
gate_compress: torch.Tensor | None = None,
freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
attention_mask: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
"""Forward pass for distributed attention.
@@ -158,6 +188,7 @@ class DistributedAttention_VSA(DistributedAttention):
replicated_q (Optional[torch.Tensor]): Replicated query tensor, typically for text tokens
replicated_k (Optional[torch.Tensor]): Replicated key tensor
replicated_v (Optional[torch.Tensor]): Replicated value tensor
attention_mask (Optional[torch.Tensor]): Attention mask [batch_size, seq_len]
Returns:
Tuple[torch.Tensor, Optional[torch.Tensor]]: A tuple containing:
@@ -173,15 +204,32 @@ class DistributedAttention_VSA(DistributedAttention):
forward_context: ForwardContext = get_forward_context()
ctx_attn_metadata = forward_context.attn_metadata
batch_size, seq_len, num_heads, head_dim = q.shape
# Stack QKV
qkvg = torch.cat([q, k, v, gate_compress],
dim=0) # [3, seq_len, num_heads, head_dim]
dim=0) # [4*batch, seq_len, num_heads, head_dim]
# Redistribute heads across sequence dimension
# Before: [4*batch, shard_seq_len, num_heads, head_dim]
# After: [4*batch, full_seq_len, shard_num_heads, head_dim]
qkvg = sequence_model_parallel_all_to_all_4D(qkvg,
scatter_dim=2,
gather_dim=1)
# After all-to-all, each rank has the full sequence but only a subset of heads
# The attention mask should now apply to the full sequence length
if attention_mask is not None:
valid_seq_len = (attention_mask[0] == 1).sum().item()
qkvg = qkvg[:, :valid_seq_len, :, :]
if freqs_cis is not None:
cos, sin = freqs_cis
qkvg[:batch_size * 2] = _apply_rotary_emb(qkvg[:batch_size * 2],
cos,
sin,
is_neox_style=False)
qkvg = self.attn_impl.preprocess_qkv(qkvg, ctx_attn_metadata)
q, k, v, gate_compress = qkvg.chunk(4, dim=0)
@@ -194,6 +242,10 @@ class DistributedAttention_VSA(DistributedAttention):
# Apply backend-specific postprocess_output
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
if attention_mask is not None:
pad_len = (attention_mask[0] == 0).sum().item()
output = torch.nn.functional.pad(output, (0, 0, 0, 0, 0, pad_len))
output = sequence_model_parallel_all_to_all_4D(output,
scatter_dim=1,
gather_dim=2)
@@ -244,6 +296,7 @@ class LocalAttention(nn.Module):
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor] | None = None,
) -> torch.Tensor:
"""
Apply local attention between query, key and value tensors.
@@ -263,5 +316,10 @@ class LocalAttention(nn.Module):
forward_context: ForwardContext = get_forward_context()
ctx_attn_metadata = forward_context.attn_metadata
if freqs_cis is not None:
cos, sin = freqs_cis
q = _apply_rotary_emb(q, cos, sin, is_neox_style=False)
k = _apply_rotary_emb(k, cos, sin, is_neox_style=False)
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
return output
@@ -0,0 +1,99 @@
# Licensed under the TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5/blob/main/LICENSE
#
# Unless and only to the extent required by applicable law, the Tencent Hunyuan works and any
# output and results there from are provided "AS IS" without any express or implied warranties of
# any kind including any warranties of title, merchantability, noninfringement, course of dealing,
# usage of trade, or fitness for a particular purpose. You are solely responsible for determining the
# appropriateness of using, reproducing, modifying, performing, displaying or distributing any of
# the Tencent Hunyuan works or outputs and assume any and all risks associated with your or a
# third party's use or distribution of any of the Tencent Hunyuan works or outputs and your exercise
# of rights and permissions under this agreement.
# See the License for the specific language governing permissions and limitations under the License.
from einops import rearrange
def flash_attn_no_pad(qkv,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False):
from flash_attn import flash_attn_varlen_qkvpacked_func
from flash_attn.bert_padding import pad_input, unpad_input
batch_size = qkv.shape[0]
seqlen = qkv.shape[1]
nheads = qkv.shape[-2]
x = rearrange(qkv, "b s three h d -> b s (three h d)")
x_unpad, indices, cu_seqlens, max_s, used_seqlens_in_batch = unpad_input(
x, key_padding_mask)
x_unpad = rearrange(x_unpad,
"nnz (three h d) -> nnz three h d",
three=3,
h=nheads)
output_unpad = flash_attn_varlen_qkvpacked_func(
x_unpad,
cu_seqlens,
max_s,
dropout_p,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic,
)
output = rearrange(
pad_input(rearrange(output_unpad, "nnz h d -> nnz (h d)"), indices,
batch_size, seqlen),
"b s (h d) -> b s h d",
h=nheads,
)
return output
def flash_attn_no_pad_v3(qkv,
key_padding_mask,
causal=False,
dropout_p=0.0,
softmax_scale=None,
deterministic=False):
from flash_attn.bert_padding import pad_input, unpad_input
from flash_attn_interface import flash_attn_varlen_func as flash_attn_varlen_func_v3
if flash_attn_varlen_func_v3 is None:
raise ImportError("FlashAttention V3 backend not available")
batch_size, seqlen, _, nheads, head_dim = qkv.shape
query, key, value = qkv.unbind(dim=2)
query_unpad, indices, cu_seqlens_q, max_seqlen_q, _ = unpad_input(
rearrange(query, "b s h d -> b s (h d)"), key_padding_mask)
key_unpad, _, cu_seqlens_k, _, _ = unpad_input(
rearrange(key, "b s h d -> b s (h d)"), key_padding_mask)
value_unpad, _, _, _, _ = unpad_input(
rearrange(value, "b s h d -> b s (h d)"), key_padding_mask)
query_unpad = rearrange(query_unpad, "nnz (h d) -> nnz h d", h=nheads)
key_unpad = rearrange(key_unpad, "nnz (h d) -> nnz h d", h=nheads)
value_unpad = rearrange(value_unpad, "nnz (h d) -> nnz h d", h=nheads)
output_unpad = flash_attn_varlen_func_v3(query_unpad,
key_unpad,
value_unpad,
cu_seqlens_q,
cu_seqlens_k,
max_seqlen_q,
max_seqlen_q,
softmax_scale=softmax_scale,
causal=causal,
deterministic=deterministic)
output = rearrange(pad_input(
rearrange(output_unpad, "nnz h d -> nnz (h d)"), indices, batch_size,
seqlen),
"b s (h d) -> b s h d",
h=nheads)
return output
+3 -2
View File
@@ -1,10 +1,11 @@
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
__all__ = [
"HunyuanVideoConfig", "WanVideoConfig", "StepVideoConfig",
"CosmosVideoConfig", "Cosmos25VideoConfig"
"HunyuanVideoConfig", "HunyuanVideo15Config", "WanVideoConfig",
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig"
]
+2
View File
@@ -23,6 +23,8 @@ class DiTArchConfig(ArchConfig):
hidden_size: int = 0
num_attention_heads: int = 0
num_channels_latents: int = 0
in_channels: int = 0
out_channels: int = 0
exclude_lora_layers: list[str] = field(default_factory=list)
boundary_ratio: float | None = None
@@ -0,0 +1,157 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_double_block(n: str, m) -> bool:
return "double_blocks" in n and str.isdigit(n.split(".")[-1])
def is_refiner_block(n: str, m) -> bool:
return "refiner" in n and str.isdigit(n.split(".")[-1])
def is_txt_in(n: str, m) -> bool:
return n.split(".")[-1] == "txt_in"
@dataclass
class HunyuanVideo15ArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(
default_factory=lambda: [is_double_block, is_refiner_block])
_compile_conditions: list = field(
default_factory=lambda: [is_double_block, is_refiner_block, is_txt_in])
param_names_mapping: dict = field(
default_factory=lambda: {
# 1. context_embedder.time_text_embed submodules (specific rules, applied first):
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_1\.(.*)$":
r"txt_in.t_embedder.mlp.fc_in.\1",
r"^context_embedder\.time_text_embed\.timestep_embedder\.linear_2\.(.*)$":
r"txt_in.t_embedder.mlp.fc_out.\1",
r"^context_embedder\.proj_in\.(.*)$":
r"txt_in.input_embedder.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_1\.(.*)$":
r"txt_in.c_embedder.fc_in.\1",
r"^context_embedder\.time_text_embed\.text_embedder\.linear_2\.(.*)$":
r"txt_in.c_embedder.fc_out.\1",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm1\.(.*)$":
r"txt_in.refiner_blocks.\1.norm1.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm2\.(.*)$":
r"txt_in.refiner_blocks.\1.norm2.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 0, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 1, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"txt_in.refiner_blocks.\1.self_attn_qkv.\2", 2, 3),
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^context_embedder\.token_refiner\.refiner_blocks\.(\d+)\.norm_out\.linear\.(.*)$":
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
# 2. txt_in_2 mapping:
r"^context_embedder_2\.(.*)$":
r"txt_in_2.\1",
# 3. x_embedder mapping:
r"^x_embedder\.proj\.(.*)$":
r"img_in.proj.\1",
# 4. Top-level time_text_embed mappings:
r"^time_embed\.timestep_embedder\.linear_1\.(.*)$":
r"time_in.timestep_embedder.mlp.fc_in.\1",
r"^time_embed\.timestep_embedder\.linear_2\.(.*)$":
r"time_in.timestep_embedder.mlp.fc_out.\1",
r"^time_embed\.timestep_embedder_r\.linear_1\.(.*)$":
r"time_in.timestep_embedder_r.mlp.fc_in.\1",
r"^time_embed\.timestep_embedder_r\.linear_2\.(.*)$":
r"time_in.timestep_embedder_r.mlp.fc_out.\1",
# 5. transformer_blocks mapping:
r"^transformer_blocks\.(\d+)\.norm1\.linear\.(.*)$":
r"double_blocks.\1.img_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.norm1_context\.linear\.(.*)$":
r"double_blocks.\1.txt_mod.linear.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_q\.(.*)$":
r"double_blocks.\1.img_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_k\.(.*)$":
r"double_blocks.\1.img_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.to_q\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_k\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_v\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_q_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_k_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
r"^transformer_blocks\.(\d+)\.attn\.add_v_proj\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
r"^transformer_blocks\.(\d+)\.attn\.to_out\.0\.(.*)$":
r"double_blocks.\1.img_attn_proj.\2",
# Corrected: merge attn.to_add_out into the main projection.
r"^transformer_blocks\.(\d+)\.attn\.to_add_out\.(.*)$":
r"double_blocks.\1.txt_attn_proj.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_q\.(.*)$":
r"double_blocks.\1.txt_attn_q_norm.\2",
r"^transformer_blocks\.(\d+)\.attn\.norm_added_k\.(.*)$":
r"double_blocks.\1.txt_attn_k_norm.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.img_mlp.fc_out.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.0(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_in.\2",
r"^transformer_blocks\.(\d+)\.ff_context\.net\.2(?:\.proj)?\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_out.\2",
# 7. Final layers mapping:
r"^norm_out\.linear\.(.*)$":
r"final_layer.adaLN_modulation.linear.\1",
r"^proj_out\.(.*)$":
r"final_layer.linear.\1",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
in_channels: int = 65
out_channels: int = 32
num_attention_heads: int = 16
attention_head_dim: int = 128
num_layers: int = 54
num_refiner_layers: int = 2
mlp_ratio: float = 4.0
patch_size: int = 1
patch_size_t: int = 1
qk_norm: str = "rms_norm"
text_embed_dim: int = 3584
text_embed_2_dim: int = 1472
image_embed_dim: int = 1152
rope_theta: float = 256.0
rope_axes_dim: tuple[int, ...] = (16, 56, 56)
target_size: int = 640
task_type: str = "i2v"
use_meanflow: bool = False
exclude_lora_layers: list[str] = field(
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
def __post_init__(self):
super().__post_init__()
self.hidden_size: int = self.attention_head_dim * self.num_attention_heads
self.num_channels_latents: int = self.out_channels
@dataclass
class HunyuanVideo15Config(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=HunyuanVideo15ArchConfig)
prefix: str = "Hunyuan15"
@@ -0,0 +1,83 @@
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig, WanVideoConfig
@dataclass
class MatrixGameWanVideoArchConfig(WanVideoArchConfig):
# Override param_names_mapping to remove patch_embedding transformation
# because MatrixGame checkpoints already have patch_embedding.proj format
param_names_mapping: dict = field(
default_factory=lambda: {
# Removed: r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1"
# because checkpoint already has correct format
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
r"condition_embedder.text_embedder.fc_in.\1",
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
r"condition_embedder.text_embedder.fc_out.\1",
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^condition_embedder\.time_proj\.(.*)$":
r"condition_embedder.time_modulation.linear.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_in.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_out.\1",
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$":
r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$":
r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$":
r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$":
r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$":
r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$":
r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
r"blocks.\1.ffn.fc_out.\2",
r"^blocks\.(\d+)\.norm2\.(.*)$":
r"blocks.\1.self_attn_residual_norm.norm.\2",
})
action_config: dict = field(
default_factory=lambda: {
"blocks": list(range(15)),
"enable_mouse": True,
"enable_keyboard": True,
"heads_num": 16,
"hidden_size": 128,
"img_hidden_size": 1536,
"keyboard_dim_in": 4,
"keyboard_hidden_dim": 1024,
"mouse_dim_in": 2,
"mouse_hidden_dim": 1024,
"mouse_qk_dim_list": [8, 28, 28],
"patch_size": [1, 2, 2],
"qk_norm": True,
"qkv_bias": False,
"rope_dim_list": [8, 28, 28],
"rope_theta": 256,
"vae_time_compression_ratio": 4,
"windows_size": 3,
})
local_attn_size: int = -1
sink_size: int = 0
num_frames_per_block: int = 3
text_len: int = 512
text_dim: int = 0
image_dim: int = 1280
@dataclass
class MatrixGameWanVideoConfig(WanVideoConfig):
arch_config: MatrixGameWanVideoArchConfig = field(
default_factory=MatrixGameWanVideoArchConfig)
prefix: str = "Wan"
@@ -6,9 +6,11 @@ from fastvideo.configs.models.encoders.clip import (
CLIPTextConfig, CLIPVisionConfig, WAN2_1ControlCLIPVisionConfig)
from fastvideo.configs.models.encoders.llama import LlamaConfig
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig",
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig"
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig",
"Qwen2_5_VLConfig"
]
@@ -72,6 +72,7 @@ class EncoderConfig(ModelConfig):
@dataclass
class TextEncoderConfig(EncoderConfig):
arch_config: ArchConfig = field(default_factory=TextEncoderArchConfig)
is_chat_model: bool = False
@dataclass
@@ -0,0 +1,93 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.encoders.base import (TextEncoderArchConfig,
TextEncoderConfig)
def _is_transformer_layer(n: str, m) -> bool:
return "layers" in n and str.isdigit(n.split(".")[-1])
def _is_embeddings(n: str, m) -> bool:
return n.endswith("embed_tokens")
def _is_final_norm(n: str, m) -> bool:
return n.endswith("norm")
@dataclass
class Qwen2_5_VLArchConfig(TextEncoderArchConfig):
vocab_size: int = 152064
hidden_size: int = 8192
intermediate_size: int = 29568
num_hidden_layers: int = 80
num_attention_heads: int = 64
num_key_value_heads: int = 8
hidden_act: str = "silu"
max_position_embeddings: int = 32768
initializer_range: float = 0.02
rms_norm_eps: float = 1e-05
use_cache: bool = True
tie_word_embeddings: bool = False
rope_theta: float = 1000000.0
use_sliding_window: bool = False
sliding_window: int | None = 4096
max_window_layers: int = 80
layer_types: list = field(default_factory=list)
attention_dropout: float = 0.0
rope_scaling: dict | None = None
bos_token_id: int | None = None
eos_token_id: int | None = None
pad_token_id: int | None = None
vision_token_id: int = 151654
model_type: str = "qwen2_5_vl_text"
dtype: str = "bfloat16"
stacked_params_mapping: list[tuple[str, str, str
| int]] = field(default_factory=lambda: [
(".qkv_proj", ".q_proj", "q"),
(".qkv_proj", ".k_proj", "k"),
(".qkv_proj", ".v_proj", "v"),
(".gate_up_proj", ".gate_proj", 0),
(".gate_up_proj", ".up_proj", 1),
])
_fsdp_shard_conditions: list = field(
default_factory=lambda:
[_is_transformer_layer, _is_embeddings, _is_final_norm])
def __post_init__(self):
super().__post_init__()
self.sliding_window = self.sliding_window if self.use_sliding_window else None
# for backward compatibility
if self.num_key_value_heads is None:
self.num_key_value_heads = self.num_attention_heads
if self.layer_types is None:
self.layer_types = [
"sliding_attention" if self.sliding_window is not None
and i >= self.max_window_layers else "full_attention"
for i in range(self.num_hidden_layers)
]
if self.rope_scaling is not None and "type" in self.rope_scaling:
if self.rope_scaling["type"] == "mrope":
self.rope_scaling["type"] = "default"
self.rope_scaling["rope_type"] = self.rope_scaling["type"]
self.tokenizer_kwargs = {
"add_generation_prompt": True,
"tokenize": True,
"return_dict": True,
"padding": "max_length",
"max_length": 1000 + 108,
"truncation": True,
"return_tensors": "pt",
}
@dataclass
class Qwen2_5_VLConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(
default_factory=Qwen2_5_VLArchConfig)
prefix: str = "qwen2_5_vl"
is_chat_model: bool = True
+3
View File
@@ -40,6 +40,8 @@ class T5ArchConfig(TextEncoderArchConfig):
eos_token_id: int = 1
classifier_dropout: float = 0.0
text_len: int = 512
dtype: str | None = None
gradient_checkpointing: bool = False
stacked_params_mapping: list[tuple[str, str,
str]] = field(default_factory=lambda: [
# (param_name, shard_name, shard_id)
@@ -68,6 +70,7 @@ class T5ArchConfig(TextEncoderArchConfig):
"return_attention_mask": True,
"return_tensors": "pt",
}
self.hidden_size = self.d_model
@dataclass
@@ -1,5 +1,6 @@
from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
from fastvideo.configs.models.vaes.stepvideovae import StepVideoVAEConfig
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
@@ -8,4 +9,5 @@ __all__ = [
"WanVAEConfig",
"StepVideoVAEConfig",
"CosmosVAEConfig",
"Hunyuan15VAEConfig",
]
@@ -0,0 +1,27 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class Hunyuan15VAEArchConfig(VAEArchConfig):
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 32
block_out_channels: tuple[int, ...] = (128, 256, 512, 1024, 1024)
layers_per_block: int = 2
spatial_compression_ratio: int = 16
temporal_compression_ratio: int = 4
downsample_match_channel: bool = True
upsample_match_channel: bool = True
scaling_factor: float = 1.03682
def __post_init__(self):
self.spatial_compression_ratio: int = 2**(len(self.block_out_channels) -
1)
@dataclass
class Hunyuan15VAEConfig(VAEConfig):
arch_config: VAEArchConfig = field(default_factory=Hunyuan15VAEArchConfig)
+5 -4
View File
@@ -2,6 +2,7 @@ from fastvideo.configs.pipelines.base import (PipelineConfig,
SlidingTileAttnConfig)
from fastvideo.configs.pipelines.cosmos import CosmosConfig
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
@@ -11,8 +12,8 @@ from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
__all__ = [
"HunyuanConfig", "FastHunyuanConfig", "PipelineConfig",
"SlidingTileAttnConfig", "WanT2V480PConfig", "WanI2V480PConfig",
"WanT2V720PConfig", "WanI2V720PConfig", "StepVideoT2VConfig",
"SelfForcingWanT2V480PConfig", "CosmosConfig",
"get_pipeline_config_cls_from_name"
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
"CosmosConfig", "get_pipeline_config_cls_from_name"
]
+139
View File
@@ -0,0 +1,139 @@
# SPDX-License-Identifier: Apache-2.0
from collections.abc import Callable
from dataclasses import dataclass, field
from typing import Any
import re
import torch
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits import HunyuanVideo15Config
from fastvideo.configs.models.encoders import (BaseEncoderOutput,
Qwen2_5_VLConfig, T5Config)
from fastvideo.configs.models.vaes import Hunyuan15VAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
PROMPT_TEMPLATE_TOKEN_LENGTH = 108
PROMPT_TEMPLATE_ENCODE_VIDEO = "You are a helpful assistant. Describe the video by detailing the following aspects: \
1. The main content and theme of the video. \
2. The color, shape, size, texture, quantity, text, and spatial relationships of the objects. \
3. Actions, events, behaviors temporal relationships, physical movement changes of the objects. \
4. background environment, light, style and atmosphere. \
5. camera angles, movements, and transitions used in the video."
def extract_glyph_texts(prompt: str) -> str | None:
"""
Extract glyph texts from prompt using regex pattern.
Args:
prompt: Input prompt string
Returns:
List of extracted glyph texts
"""
pattern = r"\"(.*?)\"|“(.*?)”"
matches = re.findall(pattern, prompt)
result = [match[0] or match[1] for match in matches]
result = list(dict.fromkeys(result)) if len(result) > 1 else result
if result:
formatted_result = ". ".join([f'Text "{text}"'
for text in result]) + ". "
else:
formatted_result = None
return formatted_result
def format_text_input(prompt: str, system_message: str) -> list[dict[str, Any]]:
"""
Apply text to template.
Args:
prompt (List[str]): Input text.
system_message (str): System message.
Returns:
List[Dict[str, Any]]: List of chat conversation.
"""
template = [{
"role": "system",
"content": system_message
}, {
"role": "user",
"content": prompt if prompt else " "
}]
return template
def qwen_preprocess_text(prompt: str) -> list[dict[str, Any]]:
output = format_text_input(prompt, PROMPT_TEMPLATE_ENCODE_VIDEO)
return output
def qwen_postprocess_text(
outputs: BaseEncoderOutput,
mask: torch.tensor) -> tuple[torch.tensor, torch.tensor]:
assert outputs.hidden_states is not None
output = outputs.hidden_states[-3]
output = output[:, PROMPT_TEMPLATE_TOKEN_LENGTH:]
mask = mask[:, PROMPT_TEMPLATE_TOKEN_LENGTH:]
return output, mask
def byt5_preprocess_text(prompt: str) -> str | None:
prompts = [prompt] if isinstance(prompt, str) else prompt
glyph_texts = [extract_glyph_texts(p) for p in prompts]
return glyph_texts[0]
def byt5_postprocess_text(outputs: BaseEncoderOutput) -> torch.tensor:
return outputs.last_hidden_state
@dataclass
class Hunyuan15T2V480PConfig(PipelineConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
# DiT
dit_config: DiTConfig = field(default_factory=HunyuanVideo15Config)
# VAE
vae_config: VAEConfig = field(default_factory=Hunyuan15VAEConfig)
# Denoising stage
flow_shift: int = 5
# Text encoding stage
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (Qwen2_5_VLConfig(), T5Config()))
preprocess_text_funcs: tuple[Callable[[Any], Any], ...] = field(
default_factory=lambda: (qwen_preprocess_text, byt5_preprocess_text))
postprocess_text_funcs: tuple[Callable[..., Any], ...] = field(
default_factory=lambda: (qwen_postprocess_text, byt5_postprocess_text))
# Precision for each component
dit_precision: str = "bf16"
vae_precision: str = "fp16"
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("bf16", "fp32"))
text_encoder_crop_start: int = PROMPT_TEMPLATE_TOKEN_LENGTH
text_encoder_max_lengths: tuple[int, ...] = field(
default_factory=lambda: (1000 + PROMPT_TEMPLATE_TOKEN_LENGTH, 256))
vae_tiling: bool = True
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
@dataclass
class Hunyuan15T2V720PConfig(Hunyuan15T2V480PConfig):
"""Base configuration for HunYuan pipeline architecture."""
# HunyuanConfig-specific parameters with defaults
flow_shift: int = 9
+32 -8
View File
@@ -7,6 +7,7 @@ from collections.abc import Callable
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.pipelines.cosmos import CosmosConfig
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
# isort: off
@@ -14,7 +15,8 @@ from fastvideo.configs.pipelines.wan import (
FastWan2_1_T2V_480P_Config, FastWan2_2_TI2V_5B_Config,
Wan2_2_I2V_A14B_Config, Wan2_2_T2V_A14B_Config, Wan2_2_TI2V_5B_Config,
WanI2V480PConfig, WanI2V720PConfig, WanT2V480PConfig, WanT2V720PConfig,
SelfForcingWanT2V480PConfig, WANV2VConfig, SelfForcingWan2_2_T2V480PConfig)
SelfForcingWanT2V480PConfig, WANV2VConfig, SelfForcingWan2_2_T2V480PConfig,
MatrixGameI2V480PConfig)
# isort: on
from fastvideo.logger import init_logger
from fastvideo.utils import (maybe_download_model_index,
@@ -26,6 +28,10 @@ logger = init_logger(__name__)
PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/FastHunyuan-diffusers": FastHunyuanConfig,
"hunyuanvideo-community/HunyuanVideo": HunyuanConfig,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
Hunyuan15T2V480PConfig,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
Hunyuan15T2V720PConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers": WANV2VConfig,
@@ -45,18 +51,32 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_Config,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_Config,
"nvidia/Cosmos-Predict2-2B-Video2World": CosmosConfig,
"FastVideo/Matrix-Game-2.0-Base-Diffusers": MatrixGameI2V480PConfig,
"FastVideo/Matrix-Game-2.0-GTA-Diffusers": MatrixGameI2V480PConfig,
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers": MatrixGameI2V480PConfig,
# Add other specific weight variants
}
# For determining pipeline type from model ID
PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
"hunyuan": lambda id: "hunyuan" in id.lower(),
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
"stepvideo": lambda id: "stepvideo" in id.lower(),
"cosmos": lambda id: "cosmos" in id.lower(),
"hunyuan":
lambda id: "hunyuan" in id.lower(),
"hunyuan15":
lambda id: "hunyuan15" in id.lower(),
"matrixgame":
lambda id: "matrix-game" in id.lower() or "matrixgame" in id.lower(),
"wanpipeline":
lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo":
lambda id: "wanimagetovideo" in id.lower(),
"wandmdpipeline":
lambda id: "wandmdpipeline" in id.lower(),
"wancausaldmdpipeline":
lambda id: "wancausaldmdpipeline" in id.lower(),
"stepvideo":
lambda id: "stepvideo" in id.lower(),
"cosmos":
lambda id: "cosmos" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -64,6 +84,9 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
"hunyuan":
HunyuanConfig, # Base Hunyuan config as fallback for any Hunyuan variant
"matrixgame": MatrixGameI2V480PConfig,
"hunyuan15":
Hunyuan15T2V480PConfig, # Base Hunyuan15 config as fallback for any Hunyuan15 variant
"wanpipeline":
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V480PConfig,
@@ -109,6 +132,7 @@ def get_pipeline_config_cls_from_name(
# First try exact match for specific weights
if pipeline_name_or_path in PIPE_NAME_TO_CONFIG:
pipeline_config_cls = PIPE_NAME_TO_CONFIG[pipeline_name_or_path]
return pipeline_config_cls
# Try partial matches (for local paths that might include the weight ID)
for registered_id, config_class in PIPE_NAME_TO_CONFIG.items():
+21
View File
@@ -6,6 +6,7 @@ import torch
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.configs.models.dits.matrixgame import MatrixGameWanVideoConfig
from fastvideo.configs.models.encoders import (BaseEncoderOutput,
CLIPVisionConfig, T5Config,
WAN2_1ControlCLIPVisionConfig)
@@ -190,3 +191,23 @@ class SelfForcingWan2_2_T2V480PConfig(Wan2_2_T2V_A14B_Config):
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
# =============================================
# ============= Matrix Game ===================
# =============================================
@dataclass
class MatrixGameI2V480PConfig(WanI2V480PConfig):
dit_config: DiTConfig = field(default_factory=MatrixGameWanVideoConfig)
image_encoder_config: EncoderConfig = field(
default_factory=WAN2_1ControlCLIPVisionConfig)
is_causal: bool = True
flow_shift: float | None = 5.0
dmd_denoising_steps: list[int] | None = field(
default_factory=lambda: [1000, 666, 333])
warp_denoising_step: bool = True
context_noise: int = 0
num_frames_per_block: int = 3
# sliding_window_num_frames: int = 15
+7
View File
@@ -17,10 +17,16 @@ class SamplingParam:
# Image inputs
image_path: str | None = None
pil_image: Any | None = None
# Video inputs
video_path: str | None = None
# Action control inputs (Matrix-Game)
mouse_cond: Any | None = None # Shape: (B, T, 2)
keyboard_cond: Any | None = None # Shape: (B, T, K)
grid_sizes: Any | None = None # Shape: (3,) [F,H,W]
# Text inputs
prompt: str | list[str] | None = None
negative_prompt: str = "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"
@@ -44,6 +50,7 @@ class SamplingParam:
guidance_scale: float = 1.0
guidance_rescale: float = 0.0
boundary_ratio: float | None = None
sigmas: list[float] | None = None
# TeaCache parameters
enable_teacache: bool = False
+29
View File
@@ -0,0 +1,29 @@
# SPDX-License-Identifier: Apache-2.0
import numpy as np
from dataclasses import dataclass, field
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class Hunyuan15_480P_SamplingParam(SamplingParam):
num_inference_steps: int = 50
num_frames: int = 121
height: int = 480
width: int = 848
fps: int = 24
guidance_scale: float = 6.0
prompt_attention_mask: list = field(default_factory=list)
negative_attention_mask: list = field(default_factory=list)
sigmas: list[float] | None = field(
default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
negative_prompt: str = ""
@dataclass
class Hunyuan15_720P_SamplingParam(Hunyuan15_480P_SamplingParam):
height: int = 720
width: int = 1280
+56 -39
View File
@@ -5,6 +5,7 @@ from typing import Any
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam)
from fastvideo.configs.sample.hunyuan15 import Hunyuan15_480P_SamplingParam, Hunyuan15_720P_SamplingParam
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
from fastvideo.configs.sample.cosmos import Cosmos_Predict2_2B_Video2World_SamplingParam
@@ -23,6 +24,7 @@ from fastvideo.configs.sample.wan import (
Wan2_1_Fun_1_3B_Control_SamplingParam,
SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
SelfForcingWan2_2_T2V_A14B_480P_SamplingParam,
MatrixGame2_SamplingParam,
)
# isort: on
from fastvideo.logger import init_logger
@@ -32,44 +34,36 @@ from fastvideo.utils import (maybe_download_model_index,
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"FastVideo/FastHunyuan-diffusers":
FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo":
HunyuanSamplingParam,
"FastVideo/stepvideo-t2v-diffusers":
StepVideoT2VSamplingParam,
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
Hunyuan15_480P_SamplingParam,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
Hunyuan15_720P_SamplingParam,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
# Wan2.1
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers":
WanT2V_1_3B_SamplingParam,
"Wan-AI/Wan2.1-T2V-14B-Diffusers":
WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers":
WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers":
WanI2V_14B_720P_SamplingParam,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers":
Wan2_1_Fun_1_3B_Control_SamplingParam,
# Wan2.2
"Wan-AI/Wan2.2-TI2V-5B-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers":
Wan2_2_T2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers":
Wan2_2_I2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_SamplingParam,
# FastWan2.1
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
FastWanT2V480P_SamplingParam,
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480P_SamplingParam,
# FastWan2.2
"FastVideo/FastWan2.2-TI2V-5B-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
# Causal Self-Forcing Wan2.1
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers":
@@ -85,17 +79,32 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"nvidia/Cosmos-Predict2-2B-Video2World":
Cosmos_Predict2_2B_Video2World_SamplingParam,
# MatrixGame2.0 models
"FastVideo/Matrix-Game-2.0-Base-Diffusers": MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-GTA-Diffusers": MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers": MatrixGame2_SamplingParam,
# Add other specific weight variants
}
# For determining pipeline type from model ID
SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
"hunyuan": lambda id: "hunyuan" in id.lower(),
"wanpipeline": lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo": lambda id: "wanimagetovideo" in id.lower(),
"stepvideo": lambda id: "stepvideo" in id.lower(),
"wandmdpipeline": lambda id: "wandmdpipeline" in id.lower(),
"wancausaldmdpipeline": lambda id: "wancausaldmdpipeline" in id.lower(),
"hunyuan":
lambda id: "hunyuan" in id.lower(),
"hunyuan15":
lambda id: "hunyuan15" in id.lower(),
"wanpipeline":
lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo":
lambda id: "wanimagetovideo" in id.lower(),
"stepvideo":
lambda id: "stepvideo" in id.lower(),
"wandmdpipeline":
lambda id: "wandmdpipeline" in id.lower(),
"wancausaldmdpipeline":
lambda id: "wancausaldmdpipeline" in id.lower(),
"matrixgame":
lambda id: "matrixgame" in id.lower() or "matrix-game" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -103,12 +112,15 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
"hunyuan":
HunyuanSamplingParam, # Base Hunyuan config as fallback for any Hunyuan variant
"hunyuan15":
Hunyuan15_480P_SamplingParam, # Base Hunyuan15 config as fallback for any Hunyuan15 variant
"wanpipeline":
WanT2V_1_3B_SamplingParam, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V_14B_480P_SamplingParam,
"wandmdpipeline": FastWanT2V480P_SamplingParam,
"wancausaldmdpipeline": SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam,
"stepvideo": StepVideoT2VSamplingParam,
"matrixgame": MatrixGame2_SamplingParam,
# Other fallbacks by architecture
}
@@ -116,6 +128,20 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
"""Get the appropriate sampling param for specific pretrained weights."""
# First try exact match for specific weights
if pipeline_name_or_path in SAMPLING_PARAM_REGISTRY:
return SAMPLING_PARAM_REGISTRY[pipeline_name_or_path]
# Try partial matches (for local paths that might include the weight ID)
for registered_id, config_class in SAMPLING_PARAM_REGISTRY.items():
if registered_id in pipeline_name_or_path:
return config_class
matrixgame_patterns = ["Matrix-Game", "Skywork--Matrix-Game", "matrixgame"]
for pattern in matrixgame_patterns:
if pattern.lower() in pipeline_name_or_path.lower():
return MatrixGame2_SamplingParam
if os.path.exists(pipeline_name_or_path):
config = verify_model_config_and_directory(pipeline_name_or_path)
logger.warning(
@@ -126,15 +152,6 @@ def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
pipeline_name = config["_class_name"]
# First try exact match for specific weights
if pipeline_name_or_path in SAMPLING_PARAM_REGISTRY:
return SAMPLING_PARAM_REGISTRY[pipeline_name_or_path]
# Try partial matches (for local paths that might include the weight ID)
for registered_id, config_class in SAMPLING_PARAM_REGISTRY.items():
if registered_id in pipeline_name_or_path:
return config_class
# If no match, try to use the fallback config
fallback_config = None
# Try to determine pipeline architecture for fallback
+11
View File
@@ -196,3 +196,14 @@ class SelfForcingWan2_2_T2V_A14B_480P_SamplingParam(
height: int = 448
width: int = 832
fps: int = 16
@dataclass
class MatrixGame2_SamplingParam(SamplingParam):
height: int = 352
width: int = 640
num_frames: int = 57
fps: int = 25
guidance_scale: float = 1.0
num_inference_steps: int = 3
negative_prompt: str | None = None
+23 -1
View File
@@ -127,7 +127,10 @@ def i2v_record_creator(batch: PreprocessBatch) -> list[dict[str, Any]]:
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]:
trajectory_timesteps: np.ndarray,
text_mask: np.ndarray | None = None,
text_embedding_2: np.ndarray | None = None,
text_mask_2: np.ndarray | None = None) -> dict[str, Any]:
"""Create a text-only ODE trajectory record matching pyarrow_schema_ode_trajectory_text_only.
Args:
@@ -165,6 +168,25 @@ def ode_text_only_record_creator(
"trajectory_timesteps_dtype": str(trajectory_timesteps.dtype),
})
if text_embedding_2 is not None:
record.update({
"text_embedding_2_bytes": text_embedding_2.tobytes(),
"text_embedding_2_shape": list(text_embedding_2.shape),
"text_embedding_2_dtype": str(text_embedding_2.dtype),
})
if text_mask is not None:
record.update({
"text_mask_bytes": text_mask.tobytes(),
"text_mask_shape": list(text_mask.shape),
"text_mask_dtype": str(text_mask.dtype),
})
if text_mask_2 is not None:
record.update({
"text_mask_2_bytes": text_mask_2.tobytes(),
"text_mask_2_shape": list(text_mask_2.shape),
"text_mask_2_dtype": str(text_mask_2.dtype),
})
return record
+9
View File
@@ -90,6 +90,15 @@ pyarrow_schema_ode_trajectory_text_only = pa.schema([
pa.field("text_embedding_shape", pa.list_(pa.int64())),
# e.g., 'bfloat16' or 'float32'
pa.field("text_embedding_dtype", pa.string()),
pa.field("text_embedding_2_bytes", pa.binary()),
pa.field("text_embedding_2_shape", pa.list_(pa.int64())),
pa.field("text_embedding_2_dtype", pa.string()),
pa.field("text_mask_bytes", pa.binary()),
pa.field("text_mask_shape", pa.list_(pa.int64())),
pa.field("text_mask_dtype", pa.string()),
pa.field("text_mask_2_bytes", pa.binary()),
pa.field("text_mask_2_shape", pa.list_(pa.int64())),
pa.field("text_mask_2_dtype", pa.string()),
# --- ODE Trajectory ---
pa.field("trajectory_latents_bytes", pa.binary()),
pa.field("trajectory_latents_shape", pa.list_(pa.int64())),
+2 -1
View File
@@ -17,6 +17,7 @@ from PIL import Image
from transformers import AutoTokenizer
from fastvideo.logger import init_logger
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
@@ -655,7 +656,7 @@ class TextDataset(torch.utils.data.IterableDataset,
self.seed = seed
# Initialize tokenizer
tokenizer_path = os.path.join(args.model_path, "tokenizer")
tokenizer_path = os.path.join(maybe_download_model(args.model_path), "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
+72 -1
View File
@@ -4,7 +4,13 @@
import torch
import torch.distributed
from fastvideo.distributed.parallel_state import get_sp_group, get_tp_group
from fastvideo.distributed.parallel_state import (get_sp_group,
get_sp_parallel_rank,
get_sp_world_size,
get_tp_group)
from fastvideo.distributed.utils import (unpad_sequence_tensor,
compute_padding_for_sp,
pad_sequence_tensor)
def tensor_model_parallel_all_reduce(input_: torch.Tensor) -> torch.Tensor:
@@ -30,3 +36,68 @@ def sequence_model_parallel_all_gather(input_: torch.Tensor,
dim: int = -1) -> torch.Tensor:
"""All-gather the input tensor across model parallel group."""
return get_sp_group().all_gather(input_, dim)
def sequence_model_parallel_all_gather_with_unpad(
input_: torch.Tensor,
original_seq_len: int,
dim: int = -1) -> torch.Tensor:
"""All-gather the input tensor and remove padding.
Args:
input_: Sharded (and possibly padded) tensor to gather
original_seq_len: Original sequence length before padding
dim: Dimension to gather along (default: -1)
Returns:
Tensor: Gathered and unpadded tensor
"""
# First gather across all ranks
gathered = get_sp_group().all_gather(input_, dim)
current_seq_len = gathered.shape[dim]
if current_seq_len > original_seq_len:
gathered = unpad_sequence_tensor(gathered,
original_seq_len,
seq_dim=dim)
return gathered
def sequence_model_parallel_shard(input_: torch.Tensor,
dim: int = 1) -> tuple[torch.Tensor, int]:
"""Shard the input tensor across model parallel group with optional padding.
Args:
input_: Input tensor to shard
dim: Dimension to shard along (default: 1)
Returns:
tuple: (sharded_tensor, original_seq_len)
- sharded_tensor: The sharded (and possibly padded) tensor
- original_seq_len: Original sequence length before padding
"""
sp_rank = get_sp_parallel_rank()
sp_world_size = get_sp_world_size()
original_seq_len = input_.shape[dim]
# Compute padding if needed
padded_seq_len, padding_amount = compute_padding_for_sp(
original_seq_len, sp_world_size)
# Pad if necessary
if padding_amount > 0:
input_ = pad_sequence_tensor(input_, padded_seq_len, seq_dim=dim)
elements_per_rank = padded_seq_len // sp_world_size
# Sharding along dim
input_ = input_.movedim(dim, 0)
input_ = input_[sp_rank * elements_per_rank:(sp_rank + 1) *
elements_per_rank]
input_ = input_.movedim(0, dim)
return input_, original_seq_len
+123
View File
@@ -61,6 +61,129 @@ def split_tensor_along_last_dim(
return tuple(tensor_list)
def compute_padding_for_sp(seq_len: int, sp_world_size: int) -> tuple[int, int]:
"""
Compute padding needed for sequence parallel.
Args:
seq_len: Original sequence length
sp_world_size: Sequence parallel world size
Returns:
tuple: (padded_seq_len, padding_amount)
"""
if seq_len % sp_world_size == 0:
return seq_len, 0
padding_amount = sp_world_size - (seq_len % sp_world_size)
padded_seq_len = seq_len + padding_amount
return padded_seq_len, padding_amount
def create_attention_mask_for_padding(
seq_len: int,
padded_seq_len: int,
batch_size: int,
device: torch.device,
dtype: torch.dtype = torch.bool,
) -> torch.Tensor | None:
"""
Create attention mask to ignore padded tokens.
Args:
seq_len: Original sequence length (before padding)
padded_seq_len: Padded sequence length
batch_size: Batch size
device: Device to create mask on
dtype: Data type for the mask (default: bool)
Returns:
Tensor: Boolean mask [B, padded_seq_len] where True = valid token,
or None if no padding is needed
"""
if seq_len == padded_seq_len:
return None
# Create mask: True for valid tokens, False for padding
attention_mask = torch.ones(
(batch_size, padded_seq_len),
dtype=dtype,
device=device,
)
# Mask out padding tokens
attention_mask[:, seq_len:] = 0
return attention_mask
def pad_sequence_tensor(
tensor: torch.Tensor,
target_seq_len: int,
seq_dim: int = 1,
pad_value: float = 0.0,
) -> torch.Tensor:
"""
Pad a tensor along the sequence dimension.
Args:
tensor: Input tensor to pad
target_seq_len: Target sequence length after padding
seq_dim: Dimension to pad along (default: 1)
pad_value: Value to use for padding (default: 0.0)
Returns:
Tensor: Padded tensor
"""
current_seq_len = tensor.shape[seq_dim]
if current_seq_len >= target_seq_len:
return tensor
padding_amount = target_seq_len - current_seq_len
# Create padding shape
pad_shape = list(tensor.shape)
pad_shape[seq_dim] = padding_amount
# Create padding tensor
padding = torch.full(
pad_shape,
pad_value,
dtype=tensor.dtype,
device=tensor.device,
)
# Concatenate along sequence dimension
padded_tensor = torch.cat([tensor, padding], dim=seq_dim)
return padded_tensor
def unpad_sequence_tensor(
tensor: torch.Tensor,
original_seq_len: int,
seq_dim: int = 1,
) -> torch.Tensor:
"""
Remove padding from a tensor along the sequence dimension.
Args:
tensor: Padded tensor
original_seq_len: Original sequence length (before padding)
seq_dim: Dimension to unpad along (default: 1)
Returns:
Tensor: Unpadded tensor
"""
# Use slice to remove padding
indices = [slice(None)] * tensor.dim()
indices[seq_dim] = slice(0, original_seq_len)
return tensor[tuple(indices)]
@dataclasses.dataclass
class StatelessProcessGroup:
"""A dataclass to hold a metadata store, and the rank, world_size of the
+18 -12
View File
@@ -102,6 +102,11 @@ class VideoGenerator:
self,
prompt: str | None = None,
sampling_param: SamplingParam | None = None,
# Action control inputs (Matrix-Game)
mouse_cond: torch.Tensor | None = None,
keyboard_cond: torch.Tensor | None = None,
grid_sizes: tuple[int, int, int] | list[int] | torch.Tensor
| None = None,
**kwargs,
) -> dict[str, Any] | list[np.ndarray] | list[dict[str, Any]]:
"""
@@ -131,6 +136,15 @@ class VideoGenerator:
if sampling_param is None:
sampling_param = SamplingParam.from_pretrained(
self.fastvideo_args.model_path)
# Add action control inputs to kwargs if provided
if mouse_cond is not None:
kwargs['mouse_cond'] = mouse_cond
if keyboard_cond is not None:
kwargs['keyboard_cond'] = keyboard_cond
if grid_sizes is not None:
kwargs['grid_sizes'] = grid_sizes
sampling_param.update(kwargs)
if self.fastvideo_args.prompt_txt is not None or sampling_param.prompt_path is not None:
@@ -297,27 +311,19 @@ class VideoGenerator:
orig_latent_num_frames = sampling_param.num_frames // 17 * 3
if orig_latent_num_frames % fastvideo_args.num_gpus != 0:
# Adjust latent frames to be divisible by number of GPUs
if sampling_param.num_frames_round_down:
# Ensure we have at least 1 batch per GPU
new_latent_num_frames = max(
1, (orig_latent_num_frames // num_gpus)) * num_gpus
else:
new_latent_num_frames = math.ceil(
orig_latent_num_frames / num_gpus) * num_gpus
if use_temporal_scaling_frames:
# Convert back to number of frames, ensuring num_frames-1 is a multiple of temporal_scale_factor
new_num_frames = (new_latent_num_frames -
new_num_frames = (orig_latent_num_frames -
1) * temporal_scale_factor + 1
else: # stepvideo only
# Find the least common multiple of 3 and num_gpus
divisor = math.lcm(3, num_gpus)
# Round up to the nearest multiple of this LCM
new_latent_num_frames = (
(new_latent_num_frames + divisor - 1) // divisor) * divisor
orig_latent_num_frames = (
(orig_latent_num_frames + divisor - 1) // divisor) * divisor
# Convert back to actual frames using the StepVideo formula
new_num_frames = new_latent_num_frames // 3 * 17
new_num_frames = orig_latent_num_frames // 3 * 17
logger.info(
"Adjusting number of frames from %s to %s based on number of GPUs (%s)",
+15
View File
@@ -32,6 +32,9 @@ if TYPE_CHECKING:
FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY: bool = False
FASTVIDEO_TORCH_PROFILER_WITH_STACK: bool = True
FASTVIDEO_TORCH_PROFILER_WITH_FLOPS: bool = False
FASTVIDEO_TORCH_PROFILER_WAIT_STEPS: int = 2
FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS: int = 1
FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS: int = 2
FASTVIDEO_TORCH_PROFILE_REGIONS: str = ""
FASTVIDEO_SERVER_DEV_MODE: bool = False
FASTVIDEO_STAGE_LOGGING: bool = False
@@ -247,6 +250,18 @@ environment_variables: dict[str, Callable[[], Any]] = {
# not profile flops.
"FASTVIDEO_TORCH_PROFILER_WITH_FLOPS":
lambda: bool(os.getenv("FASTVIDEO_TORCH_PROFILER_WITH_FLOPS", "0") != "0"),
# Wait steps per profiling cycle (torch.profiler.schedule wait parameter)
# Defaults to 2 if not set.
"FASTVIDEO_TORCH_PROFILER_WAIT_STEPS":
lambda: int(os.getenv("FASTVIDEO_TORCH_PROFILER_WAIT_STEPS", "2")),
# Warmup steps per profiling cycle (torch.profiler.schedule warmup parameter)
# Defaults to 1 if not set.
"FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS":
lambda: int(os.getenv("FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS", "1")),
# Active steps per profiling cycle (torch.profiler.schedule active parameter)
# Defaults to 2 if not set.
"FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS":
lambda: int(os.getenv("FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS", "2")),
"FASTVIDEO_TORCH_PROFILE_REGIONS":
lambda: os.getenv("FASTVIDEO_TORCH_PROFILE_REGIONS", ""),
+8
View File
@@ -174,6 +174,8 @@ class FastVideoArgs:
init_weights_from_safetensors: str = "" # path to safetensors file for initial weight loading
init_weights_from_safetensors_2: str = "" # path to safetensors file for initial weight loading for transformer_2
override_pipeline_cls_name: str | None = None
# # DMD parameters
# dmd_denoising_steps: List[int] | None = field(default=None)
@@ -424,6 +426,12 @@ class FastVideoArgs:
default=FastVideoArgs.override_transformer_cls_name,
help="Override transformer cls name",
)
parser.add_argument(
"--override-pipeline-cls-name",
type=str,
default=FastVideoArgs.override_pipeline_cls_name,
help="Override pipeline cls name",
)
parser.add_argument(
"--init-weights-from-safetensors",
type=str,
+1
View File
@@ -87,6 +87,7 @@ _ACTIVATION_REGISTRY = {
"gelu_pytorch_tanh": lambda: nn.GELU(approximate="tanh"),
"relu": nn.ReLU,
"silu": nn.SiLU,
"swish": nn.SiLU,
"quick_gelu": QuickGELU,
}
+9 -3
View File
@@ -429,6 +429,7 @@ def get_rotary_pos_embed(
theta_rescale_factor=1.0,
interpolation_factor=1.0,
shard_dim: int = 0,
do_sp_sharding: bool = False,
dtype: torch.dtype = torch.float32,
start_frame: int = 0,
) -> tuple[torch.Tensor, torch.Tensor]:
@@ -444,6 +445,7 @@ def get_rotary_pos_embed(
theta_rescale_factor: Rescale factor for theta. Defaults to 1.0
interpolation_factor: Factor to scale positions. Defaults to 1.0
shard_dim: Which dimension to shard for sequence parallelism. Defaults to 0.
do_sp_sharding: Whether to shard the positional embeddings for sequence parallelism. Defaults to False.
Returns:
Tuple of (cos, sin) tensors for rotary embeddings
@@ -460,9 +462,13 @@ def get_rotary_pos_embed(
) == head_dim, "sum(rope_dim_list) should equal to head_dim of attention layer"
# Get SP info
sp_group = get_sp_group()
sp_rank = sp_group.rank_in_group
sp_world_size = sp_group.world_size
if do_sp_sharding:
sp_group = get_sp_group()
sp_rank = sp_group.rank_in_group
sp_world_size = sp_group.world_size
else:
sp_rank = 0
sp_world_size = 1
freqs_cos, freqs_sin = get_nd_rotary_pos_embed(
rope_dim_list,
+8 -19
View File
@@ -23,6 +23,7 @@ from fastvideo.layers.visual_embedding import (ModulateProjection, PatchEmbed,
from fastvideo.models.dits.base import CachableDiT
from fastvideo.models.utils import modulate
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.distributed.communication_op import sequence_model_parallel_shard, sequence_model_parallel_all_gather
class HunyuanRMSNorm(nn.Module):
@@ -239,14 +240,7 @@ class MMDoubleStreamBlock(nn.Module):
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Apply rotary embeddings
cos, sin = freqs_cis
img_q, img_k = _apply_rotary_emb(
img_q, cos, sin,
is_neox_style=False), _apply_rotary_emb(img_k,
cos,
sin,
is_neox_style=False)
# Prepare text for attention using fused operation
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
@@ -265,7 +259,7 @@ class MMDoubleStreamBlock(nn.Module):
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
# Run distributed attention
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v)
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v, freqs_cis=freqs_cis)
img_attn_out, _ = self.img_attn_proj(
img_attn.view(batch_size, image_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
@@ -395,18 +389,11 @@ class MMSingleStreamBlock(nn.Module):
img_q, txt_q = q[:, :-txt_len], q[:, -txt_len:]
img_k, txt_k = k[:, :-txt_len], k[:, -txt_len:]
img_v, txt_v = v[:, :-txt_len], v[:, -txt_len:]
# Apply rotary embeddings to image parts
cos, sin = freqs_cis
img_q, img_k = _apply_rotary_emb(
img_q, cos, sin,
is_neox_style=False), _apply_rotary_emb(img_k,
cos,
sin,
is_neox_style=False)
# Run distributed attention
img_attn_output, txt_attn_output = self.attn(img_q, img_k, img_v, txt_q,
txt_k, txt_v)
txt_k, txt_v, freqs_cis = freqs_cis)
attn_output = torch.cat((img_attn_output, txt_attn_output),
dim=1).view(batch_size, seq_len, -1)
# Process MLP activation
@@ -593,7 +580,7 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
# Get rotary embeddings
freqs_cos, freqs_sin = get_rotary_pos_embed(
(tt * get_sp_world_size(), th, tw), self.hidden_size,
(tt, th, tw), self.hidden_size,
self.num_attention_heads, self.rope_dim_list, self.rope_theta)
freqs_cos = freqs_cos.to(x.device)
freqs_sin = freqs_sin.to(x.device)
@@ -608,6 +595,7 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
vec = vec + self.guidance_in(guidance)
# Embed image and text
img = self.img_in(img)
img, _ = sequence_model_parallel_shard(img, dim=1)
txt = self.txt_in(txt, t)
txt_seq_len = txt.shape[1]
img_seq_len = img.shape[1]
@@ -648,6 +636,7 @@ class HunyuanVideoTransformer3DModel(CachableDiT):
self.maybe_cache_states(img, original_img)
# Final layer processing
img = sequence_model_parallel_all_gather(img, dim=1)
img = self.final_layer(img, vec)
# Unpatchify to get original shape
img = unpatchify(img, tt, th, tw, self.patch_size, self.out_channels)
+853
View File
@@ -0,0 +1,853 @@
# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Any, Dict, Optional, List
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.attention import DistributedAttention, LocalAttention
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather_with_unpad,
sequence_model_parallel_shard)
from fastvideo.configs.models.dits import HunyuanVideo15Config
from fastvideo.layers.layernorm import (LayerNormScaleShift, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.layers.linear import ReplicatedLinear
# TODO(will-PY-refactor): RMSNorm ....
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
from fastvideo.layers.visual_embedding import (ModulateProjection, PatchEmbed,
TimestepEmbedder, unpatchify)
from fastvideo.models.dits.base import CachableDiT
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.logger import init_logger
from fastvideo.forward_context import set_forward_context
from fastvideo.attention.backends.abstract import AttentionMetadata
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.distributed.utils import create_attention_mask_for_padding
logger = init_logger(__name__)
class HunyuanRMSNorm(nn.Module):
def __init__(
self,
dim: int,
elementwise_affine=True,
eps: float = 1e-6,
device=None,
dtype=None,
):
"""
Initialize the RMSNorm normalization layer.
Args:
dim (int): The dimension of the input tensor.
eps (float, optional): A small value added to the denominator for numerical stability. Default is 1e-6.
Attributes:
eps (float): A small value added to the denominator for numerical stability.
weight (nn.Parameter): Learnable scaling parameter.
"""
factory_kwargs = {"device": device, "dtype": dtype}
super().__init__()
self.eps = eps
if elementwise_affine:
self.weight = nn.Parameter(torch.ones(dim, **factory_kwargs))
def _norm(self, x) -> torch.Tensor:
"""
Apply the RMSNorm normalization to the input tensor.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The normalized tensor.
"""
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)
def forward(self, x):
"""
Forward pass through the RMSNorm layer.
Args:
x (torch.Tensor): The input tensor.
Returns:
torch.Tensor: The output tensor after applying RMSNorm.
"""
output = self._norm(x.float()).type_as(x)
if hasattr(self, "weight"):
output = output * self.weight
return output
class HunyuanVideo15TimeEmbedding(nn.Module):
r"""
Time embedding for HunyuanVideo 1.5.
Supports standard timestep embedding and optional reference timestep embedding for MeanFlow-based super-resolution
models.
Args:
embedding_dim (`int`):
The dimension of the output embedding.
"""
def __init__(self, embedding_dim: int, use_meanflow: bool = False):
super().__init__()
self.timestep_embedder = TimestepEmbedder(hidden_size=embedding_dim)
self.use_meanflow = use_meanflow
self.time_proj_r = None
self.timestep_embedder_r = None
if use_meanflow:
self.timestep_embedder_r = TimestepEmbedder(hidden_size=embedding_dim)
def forward(
self,
timestep: torch.Tensor,
timestep_r: Optional[torch.Tensor] = None,
) -> torch.Tensor:
timesteps_emb = self.timestep_embedder(timestep)
if timestep_r is not None:
timesteps_emb_r = self.timestep_embedder_r(timestep_r)
timesteps_emb = timesteps_emb + timesteps_emb_r
return timesteps_emb
class HunyuanVideo15ByT5TextProjection(nn.Module):
def __init__(self, in_features: int, hidden_size: int, out_features: int):
super().__init__()
self.norm = nn.LayerNorm(in_features)
self.linear_1 = nn.Linear(in_features, hidden_size)
self.linear_2 = nn.Linear(hidden_size, hidden_size)
self.linear_3 = nn.Linear(hidden_size, out_features)
self.act_fn = nn.GELU()
def forward(self, encoder_hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.norm(encoder_hidden_states)
hidden_states = self.linear_1(hidden_states)
hidden_states = self.act_fn(hidden_states)
hidden_states = self.linear_2(hidden_states)
hidden_states = self.act_fn(hidden_states)
hidden_states = self.linear_3(hidden_states)
return hidden_states
class HunyuanVideo15ImageProjection(nn.Module):
def __init__(self, in_channels: int, hidden_size: int):
super().__init__()
self.norm_in = nn.LayerNorm(in_channels)
self.linear_1 = nn.Linear(in_channels, in_channels)
self.act_fn = nn.GELU()
self.linear_2 = nn.Linear(in_channels, hidden_size)
self.norm_out = nn.LayerNorm(hidden_size)
def forward(self, image_embeds: torch.Tensor) -> torch.Tensor:
hidden_states = self.norm_in(image_embeds)
hidden_states = self.linear_1(hidden_states)
hidden_states = self.act_fn(hidden_states)
hidden_states = self.linear_2(hidden_states)
hidden_states = self.norm_out(hidden_states)
return hidden_states
class MMDoubleStreamBlock(nn.Module):
"""
A multimodal DiT block with separate modulation for text and image/video,
using distributed attention and linear layers.
"""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float,
dtype: torch.dtype | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = "",
):
super().__init__()
self.deterministic = False
self.num_attention_heads = num_attention_heads
head_dim = hidden_size // num_attention_heads
mlp_hidden_dim = int(hidden_size * mlp_ratio)
# Image modulation components
self.img_mod = ModulateProjection(
hidden_size,
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.img_mod",
)
# Fused operations for image stream
self.img_attn_norm = LayerNormScaleShift(hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.img_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.img_mlp_residual = ScaleResidual()
# Image attention components
self.img_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_qkv")
self.img_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.img_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.img_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_proj")
self.img_mlp = MLP(hidden_size,
mlp_hidden_dim,
bias=True,
dtype=dtype,
prefix=f"{prefix}.img_mlp")
# Text modulation components
self.txt_mod = ModulateProjection(
hidden_size,
factor=6,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.txt_mod",
)
# Fused operations for text stream
self.txt_attn_norm = LayerNormScaleShift(hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.txt_attn_residual_mlp_norm = ScaleResidualLayerNormScaleShift(
hidden_size,
norm_type="layer",
elementwise_affine=False,
dtype=dtype)
self.txt_mlp_residual = ScaleResidual()
# Text attention components
self.txt_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=True,
params_dtype=dtype)
# QK norm layers for text
self.txt_attn_q_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.txt_attn_k_norm = HunyuanRMSNorm(head_dim, eps=1e-6, dtype=dtype)
self.txt_attn_proj = ReplicatedLinear(hidden_size,
hidden_size,
bias=True,
params_dtype=dtype)
self.txt_mlp = MLP(hidden_size, mlp_hidden_dim, bias=True, dtype=dtype)
# Distributed attention
self.attn = DistributedAttention(
num_heads=num_attention_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn")
def forward(
self,
img: torch.Tensor,
txt: torch.Tensor,
encoder_attention_mask: torch.Tensor,
vec: torch.Tensor,
freqs_cis: tuple,
seq_attention_mask: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
# Process modulation vectors
img_mod_outputs = self.img_mod(vec)
(
img_attn_shift,
img_attn_scale,
img_attn_gate,
img_mlp_shift,
img_mlp_scale,
img_mlp_gate,
) = torch.chunk(img_mod_outputs, 6, dim=-1)
txt_mod_outputs = self.txt_mod(vec)
(
txt_attn_shift,
txt_attn_scale,
txt_attn_gate,
txt_mlp_shift,
txt_mlp_scale,
txt_mlp_gate,
) = torch.chunk(txt_mod_outputs, 6, dim=-1)
# Prepare image for attention using fused operation
img_attn_input = self.img_attn_norm(img, img_attn_shift, img_attn_scale)
# Get QKV for image
img_qkv, _ = self.img_attn_qkv(img_attn_input)
batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
# Split QKV
img_qkv = img_qkv.view(batch_size, image_seq_len, 3,
self.num_attention_heads, -1)
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :,
2]
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Prepare text for attention using fused operation
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale)
# Get QKV for text
txt_qkv, _ = self.txt_attn_qkv(txt_attn_input)
batch_size, text_seq_len = txt_qkv.shape[0], txt_qkv.shape[1]
# Split QKV
txt_qkv = txt_qkv.view(batch_size, text_seq_len, 3,
self.num_attention_heads, -1)
txt_q, txt_k, txt_v = txt_qkv[:, :, 0], txt_qkv[:, :, 1], txt_qkv[:, :,
2]
# Apply QK-Norm if needed
txt_q = self.txt_attn_q_norm(txt_q).to(txt_q.dtype)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
# seq_len = txt_q.shape[1] + img_q.shape[1]
# attention_mask = F.pad(encoder_attention_mask, (seq_len - encoder_attention_mask.shape[1], 0), value=True)
# attention_mask = attention_mask.bool()
# self_attn_mask_1 = attention_mask.view(batch_size, 1, 1, seq_len).repeat(1, 1, seq_len, 1)
# self_attn_mask_2 = self_attn_mask_1.transpose(2, 3)
# attention_mask = (self_attn_mask_1 & self_attn_mask_2).bool()
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata = FlashAttnMetadataBuilder().build(
current_timestep=0,
attn_mask=encoder_attention_mask,
)
# Run distributed attention
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v, freqs_cis=freqs_cis, attention_mask=seq_attention_mask)
img_attn_out, _ = self.img_attn_proj(
img_attn.view(batch_size, image_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
img_mlp_input, img_residual = self.img_attn_residual_mlp_norm(
img, img_attn_out, img_attn_gate, img_mlp_shift, img_mlp_scale)
# Process image MLP
img_mlp_out = self.img_mlp(img_mlp_input)
img = self.img_mlp_residual(img_residual, img_mlp_out, img_mlp_gate)
# Process text attention output
txt_attn_out, _ = self.txt_attn_proj(
txt_attn.reshape(batch_size, text_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
txt_mlp_input, txt_residual = self.txt_attn_residual_mlp_norm(
txt, txt_attn_out, txt_attn_gate, txt_mlp_shift, txt_mlp_scale)
# Process text MLP
txt_mlp_out = self.txt_mlp(txt_mlp_input)
txt = self.txt_mlp_residual(txt_residual, txt_mlp_out, txt_mlp_gate)
return img, txt
class HunyuanVideo15Transformer3DModel(CachableDiT):
r"""
A Transformer model for video-like data used in [HunyuanVideo1.5](https://huggingface.co/tencent/HunyuanVideo1.5).
"""
# shard single stream, double stream blocks, and refiner_blocks
_fsdp_shard_conditions = HunyuanVideo15Config()._fsdp_shard_conditions
_compile_conditions = HunyuanVideo15Config()._compile_conditions
_supported_attention_backends = HunyuanVideo15Config(
)._supported_attention_backends
param_names_mapping = HunyuanVideo15Config().param_names_mapping
reverse_param_names_mapping = HunyuanVideo15Config(
).reverse_param_names_mapping
lora_param_names_mapping = HunyuanVideo15Config().lora_param_names_mapping
def __init__(
self,
config: HunyuanVideo15Config,
hf_config: dict[str, Any],
) -> None:
super().__init__(config=config, hf_config=hf_config)
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.num_channels_latents = config.num_channels_latents
self.out_channels = config.out_channels or config.in_channels
self.patch_size = (config.patch_size_t, config.patch_size, config.patch_size)
# 1. Latent and condition embedders
self.img_in = PatchEmbed(self.patch_size,
config.in_channels,
self.hidden_size,
prefix=f"{config.prefix}.img_in")
self.image_embedder = HunyuanVideo15ImageProjection(config.image_embed_dim, self.hidden_size)
self.txt_in = SingleTokenRefiner(config.text_embed_dim,
self.hidden_size,
config.num_attention_heads,
depth=config.num_refiner_layers,
dtype=None,
prefix=f"{config.prefix}.txt_in")
self.txt_in_2 = HunyuanVideo15ByT5TextProjection(config.text_embed_2_dim, 2048, self.hidden_size)
self.time_in = HunyuanVideo15TimeEmbedding(self.hidden_size, use_meanflow=config.use_meanflow)
self.cond_type_embed = nn.Embedding(3, self.hidden_size)
# 3. Dual stream transformer blocks
self.double_blocks = nn.ModuleList(
[
MMDoubleStreamBlock(
hidden_size=self.hidden_size,
num_attention_heads=config.num_attention_heads,
mlp_ratio=config.mlp_ratio,
dtype=None,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{config.prefix}.double_blocks.{i}"
)
for i in range(config.num_layers)
]
)
# 5. Output projection
self.final_layer = FinalLayer(self.hidden_size,
self.patch_size,
self.out_channels,
prefix=f"{config.prefix}.final_layer")
self.gradient_checkpointing = False
self.__post_init__()
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: List[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: List[torch.Tensor],
encoder_attention_mask: List[torch.Tensor],
guidance: Optional[torch.Tensor] = None,
timestep_r: Optional[torch.LongTensor] = None,
attention_kwargs: Optional[Dict[str, Any]] = None,
):
encoder_hidden_states_image = encoder_hidden_states_image[0]
encoder_hidden_states, encoder_hidden_states_2 = encoder_hidden_states
encoder_attention_mask, encoder_attention_mask_2 = encoder_attention_mask
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.config.patch_size_t, self.config.patch_size, self.config.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# 1. RoPE
# Get rotary embeddings
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames, post_patch_height, post_patch_width), self.hidden_size,
self.num_attention_heads, self.config.rope_axes_dim, self.config.rope_theta)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# 2. Conditional embeddings
temb = self.time_in(timestep, timestep_r=timestep_r)
hidden_states = self.img_in(hidden_states)
hidden_states, original_seq_len = sequence_model_parallel_shard(hidden_states, dim=1)
current_seq_len = hidden_states.shape[1]
sp_world_size = get_sp_world_size()
padded_seq_len = current_seq_len * sp_world_size
if padded_seq_len > original_seq_len:
seq_attention_mask = create_attention_mask_for_padding(
seq_len=original_seq_len,
padded_seq_len=padded_seq_len,
batch_size=batch_size,
device=hidden_states.device,
)
else:
seq_attention_mask = None
# qwen text embedding
encoder_hidden_states = self.txt_in(encoder_hidden_states, timestep, encoder_attention_mask)
encoder_hidden_states_cond_emb = self.cond_type_embed(
torch.zeros_like(encoder_hidden_states[:, :, 0], dtype=torch.long)
)
encoder_hidden_states = encoder_hidden_states + encoder_hidden_states_cond_emb
# byt5 text embedding
encoder_hidden_states_2 = self.txt_in_2(encoder_hidden_states_2)
encoder_hidden_states_2_cond_emb = self.cond_type_embed(
torch.ones_like(encoder_hidden_states_2[:, :, 0], dtype=torch.long)
)
encoder_hidden_states_2 = encoder_hidden_states_2 + encoder_hidden_states_2_cond_emb
# image embed
encoder_hidden_states_3 = self.image_embedder(encoder_hidden_states_image)
is_t2v = torch.all(encoder_hidden_states_image == 0)
if is_t2v:
encoder_hidden_states_3 = encoder_hidden_states_3 * 0.0
encoder_attention_mask_3 = torch.zeros(
(batch_size, encoder_hidden_states_3.shape[1]),
dtype=encoder_attention_mask.dtype,
device=encoder_attention_mask.device,
)
else:
encoder_attention_mask_3 = torch.ones(
(batch_size, encoder_hidden_states_3.shape[1]),
dtype=encoder_attention_mask.dtype,
device=encoder_attention_mask.device,
)
encoder_hidden_states_3_cond_emb = self.cond_type_embed(
2
* torch.ones_like(
encoder_hidden_states_3[:, :, 0],
dtype=torch.long,
)
)
encoder_hidden_states_3 = encoder_hidden_states_3 + encoder_hidden_states_3_cond_emb
# reorder and combine text tokens: combine valid tokens first, then padding
encoder_attention_mask = encoder_attention_mask.bool()
encoder_attention_mask_2 = encoder_attention_mask_2.bool()
encoder_attention_mask_3 = encoder_attention_mask_3.bool()
new_encoder_hidden_states = []
new_encoder_attention_mask = []
for text, text_mask, text_2, text_mask_2, image, image_mask in zip(
encoder_hidden_states,
encoder_attention_mask,
encoder_hidden_states_2,
encoder_attention_mask_2,
encoder_hidden_states_3,
encoder_attention_mask_3,
):
# Concatenate: [valid_image, valid_byt5, valid_mllm, invalid_image, invalid_byt5, invalid_mllm]
new_encoder_hidden_states.append(
torch.cat(
[
image[image_mask], # valid image
text_2[text_mask_2], # valid byt5
text[text_mask], # valid mllm
image[~image_mask], # invalid image (zeroed)
torch.zeros_like(text_2[~text_mask_2]), # invalid byt5 (zeroed)
torch.zeros_like(text[~text_mask]), # invalid mllm (zeroed)
],
dim=0,
)
)
# Apply same reordering to attention masks
new_encoder_attention_mask.append(
torch.cat(
[
image_mask[image_mask],
text_mask_2[text_mask_2],
text_mask[text_mask],
image_mask[~image_mask],
text_mask_2[~text_mask_2],
text_mask[~text_mask],
],
dim=0,
)
)
encoder_hidden_states = torch.stack(new_encoder_hidden_states)
encoder_attention_mask = torch.stack(new_encoder_attention_mask)
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.double_blocks:
hidden_states, encoder_hidden_states = self._gradient_checkpointing_func(
block,
hidden_states,
encoder_hidden_states,
encoder_attention_mask,
temb,
freqs_cis,
seq_attention_mask
)
else:
for block in self.double_blocks:
hidden_states, encoder_hidden_states = block(
hidden_states,
encoder_hidden_states,
encoder_attention_mask,
temb,
freqs_cis,
seq_attention_mask
)
# Final layer processing
hidden_states = sequence_model_parallel_all_gather_with_unpad(hidden_states, original_seq_len, dim=1)
hidden_states = self.final_layer(hidden_states, temb)
# Unpatchify to get original shape
hidden_states = unpatchify(hidden_states, post_patch_num_frames, post_patch_height, post_patch_width, self.patch_size, self.out_channels)
return hidden_states
class SingleTokenRefiner(nn.Module):
"""
A token refiner that processes text embeddings with attention to improve
their representation for cross-attention with image features.
"""
def __init__(
self,
in_channels,
hidden_size,
num_attention_heads,
depth=2,
qkv_bias=True,
dtype=None,
prefix: str = "",
) -> None:
super().__init__()
# Input projection
# self.input_embedder = ReplicatedLinear(
# in_channels,
# hidden_size,
# bias=True,
# params_dtype=dtype,
# prefix=f"{prefix}.input_embedder")
self.input_embedder = nn.Linear(in_channels, hidden_size, bias=True)
# Timestep embedding
self.t_embedder = TimestepEmbedder(hidden_size,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.t_embedder")
# Context embedding
self.c_embedder = MLP(in_channels,
hidden_size,
hidden_size,
act_type="silu",
dtype=dtype,
prefix=f"{prefix}.c_embedder")
# Refiner blocks
self.refiner_blocks = nn.ModuleList([
IndividualTokenRefinerBlock(
hidden_size,
num_attention_heads,
qkv_bias=qkv_bias,
dtype=dtype,
prefix=f"{prefix}.refiner_blocks.{i}",
) for i in range(depth)
])
def forward(self, x, t, mask=None):
# Get timestep embeddings
timestep_aware_representations = self.t_embedder(t)
# Get context-aware representations
original_dtype = x.dtype
if mask is None:
context_aware_representations = x.mean(dim=1)
else:
mask_float = mask.float().unsqueeze(-1) # [B, L, 1]
context_aware_representations = (x * mask_float).sum(dim=1) / mask_float.sum(dim=1)
context_aware_representations = self.c_embedder(
context_aware_representations)
c = timestep_aware_representations + context_aware_representations
# Project input
x = self.input_embedder(x)
# Process through refiner blocks
for block in self.refiner_blocks:
x = block(x, c, mask)
return x
class IndividualTokenRefinerBlock(nn.Module):
"""
A transformer block for refining individual tokens with self-attention.
"""
def __init__(
self,
hidden_size,
num_attention_heads,
mlp_ratio=4.0,
qkv_bias=True,
dtype=None,
prefix: str = "",
) -> None:
super().__init__()
self.num_attention_heads = num_attention_heads
mlp_hidden_dim = int(hidden_size * mlp_ratio)
# Normalization and attention
self.norm1 = nn.LayerNorm(hidden_size,
eps=1e-6,
elementwise_affine=True,
dtype=dtype)
self.self_attn_qkv = ReplicatedLinear(hidden_size,
hidden_size * 3,
bias=qkv_bias,
params_dtype=dtype,
prefix=f"{prefix}.self_attn_qkv")
self.self_attn_proj = ReplicatedLinear(
hidden_size,
hidden_size,
bias=qkv_bias,
params_dtype=dtype,
prefix=f"{prefix}.self_attn_proj")
# MLP
self.norm2 = nn.LayerNorm(hidden_size,
eps=1e-6,
elementwise_affine=True,
dtype=dtype)
self.mlp = MLP(hidden_size,
mlp_hidden_dim,
bias=True,
act_type="silu",
dtype=dtype,
prefix=f"{prefix}.mlp")
# Modulation
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.adaLN_modulation")
# Scaled dot product attention
self.attn = LocalAttention(
num_heads=num_attention_heads,
head_size=hidden_size // num_attention_heads,
# TODO: remove hardcode; remove STA
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA),
)
def forward(self, x, c, mask=None):
if mask is not None:
mask = mask.clone().bool()
mask[:, 0] = True # Prevent attention weights from becoming NaN
# Get modulation parameters
gate_msa, gate_mlp = self.adaLN_modulation(c).chunk(2, dim=-1)
# Self-attention
norm_x = self.norm1(x)
qkv, _ = self.self_attn_qkv(norm_x)
batch_size, seq_len = qkv.shape[0], qkv.shape[1]
qkv = qkv.view(batch_size, seq_len, 3, self.num_attention_heads, -1)
q, k, v = qkv[:, :, 0], qkv[:, :, 1], qkv[:, :, 2]
# Run scaled dot product attention
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata = FlashAttnMetadataBuilder().build(
current_timestep=0,
attn_mask=mask,
)
# Run distributed attention
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
attn_output = self.attn(q, k, v) # [B, L, H, D]
attn_output = attn_output.reshape(batch_size, seq_len,
-1) # [B, L, H*D]
# Project and apply residual connection with gating
attn_out, _ = self.self_attn_proj(attn_output)
x = x + attn_out * gate_msa.unsqueeze(1)
# MLP
mlp_out = self.mlp(self.norm2(x))
x = x + mlp_out * gate_mlp.unsqueeze(1)
return x
class FinalLayer(nn.Module):
"""
The final layer of DiT that projects features to pixel space.
"""
def __init__(self,
hidden_size,
patch_size,
out_channels,
dtype=None,
prefix: str = "") -> None:
super().__init__()
# Normalization
self.norm_final = nn.LayerNorm(hidden_size,
eps=1e-6,
elementwise_affine=False,
dtype=dtype)
output_dim = patch_size[0] * patch_size[1] * patch_size[2] * out_channels
self.linear = ReplicatedLinear(hidden_size,
output_dim,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.linear")
# Modulation
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.adaLN_modulation")
def forward(self, x, c):
# What the heck HF? Why you change the scale and shift order here???
scale, shift = self.adaLN_modulation(c).chunk(2, dim=-1)
x = self.norm_final(x) * (1.0 + scale.unsqueeze(1)) + shift.unsqueeze(1)
x, _ = self.linear(x)
return x
@@ -0,0 +1,11 @@
from .model import MatrixGameWanModel, MatrixGameTransformerBlock
from .causal_model import CausalMatrixGameWanModel, CausalMatrixGameTransformerBlock
from .action_module import ActionModule
__all__ = [
"MatrixGameWanModel",
"MatrixGameTransformerBlock",
"CausalMatrixGameWanModel",
"CausalMatrixGameTransformerBlock",
"ActionModule",
]
@@ -0,0 +1,567 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from Matrix-Game: https://github.com/SkyworkAI/Matrix-Game/blob/main/Matrix-Game-2/wan/modules/action_module.py
from einops import rearrange
import torch
import torch.nn as nn
import math
from torch.nn.attention.flex_attention import flex_attention
from fastvideo.attention import LocalAttention
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.layernorm import FP32LayerNorm, RMSNorm
from fastvideo.layers.rotary_embedding import (
get_nd_rotary_pos_embed as _fv_get_nd_rotary_pos_embed,
_apply_rotary_emb,
)
from fastvideo.platforms import AttentionBackendEnum
DISABLE_COMPILE = False
flex_attention = torch.compile(
flex_attention, dynamic=False, mode="max-autotune-no-cudagraphs")
def _get_nd_rotary_pos_embed_matrixgame(
rope_dim_list,
rope_sizes,
theta: float = 10000.0,
theta_rescale_factor: float = 1.0,
):
cos, sin = _fv_get_nd_rotary_pos_embed(
rope_dim_list,
rope_sizes,
theta=theta,
theta_rescale_factor=theta_rescale_factor,
dtype=torch.float32,
)
# convert from [S, D/2] to [S, D] format
cos = cos.repeat_interleave(2, dim=1)
sin = sin.repeat_interleave(2, dim=1)
return cos, sin
def _apply_rotary_emb_qk(
xq: torch.Tensor,
xk: torch.Tensor,
freqs_cos: torch.Tensor,
freqs_sin: torch.Tensor,
start_offset: int = 0,
) -> tuple[torch.Tensor, torch.Tensor]:
seq_len = xq.shape[1]
# Slice frequencies based on offset
cos = freqs_cos[start_offset:start_offset + seq_len] # [S, D]
sin = freqs_sin[start_offset:start_offset + seq_len] # [S, D]
# Move to device
cos = cos.to(xq.device)
sin = sin.to(xq.device)
# Convert from [S, D] (interleaved) back to [S, D/2]
cos_half = cos[:, ::2] # [S, D/2]
sin_half = sin[:, ::2] # [S, D/2]
# xq/xk are [B, S, H, D], need to reshape for each batch
B, S, H, D = xq.shape
xq_out = _apply_rotary_emb(xq, cos_half, sin_half, is_neox_style=False)
xk_out = _apply_rotary_emb(xk, cos_half, sin_half, is_neox_style=False)
return xq_out, xk_out
class ActionModule(nn.Module):
"""
action module from https://arxiv.org/pdf/2501.08325
"""
def __init__(
self,
mouse_dim_in: int = 2,
keyboard_dim_in: int = 6,
hidden_size: int = 128,
img_hidden_size: int = 1536,
keyboard_hidden_dim: int = 1024,
mouse_hidden_dim: int = 1024,
vae_time_compression_ratio: int = 4,
windows_size: int = 3,
heads_num: int = 16,
patch_size: list | None = None,
qk_norm: bool = True,
qkv_bias: bool = False,
rope_dim_list: list | None = None,
rope_theta = 256,
mouse_qk_dim_list: list | None = None,
enable_mouse = True,
enable_keyboard = True,
local_attn_size = 6,
blocks: list | None = None,
):
super().__init__()
# Initialize mutable defaults
patch_size = patch_size if patch_size is not None else [1, 2, 2]
rope_dim_list = rope_dim_list if rope_dim_list is not None else [8, 28, 28]
mouse_qk_dim_list = mouse_qk_dim_list if mouse_qk_dim_list is not None else [8, 28, 28]
blocks = blocks if blocks is not None else []
self.local_attn_size = local_attn_size
self.enable_mouse = enable_mouse
self.enable_keyboard = enable_keyboard
self.rope_dim_list = rope_dim_list
self.rope_theta = rope_theta
if self.enable_keyboard:
self.keyboard_embed = nn.Sequential(
nn.Linear(keyboard_dim_in, hidden_size, bias=True),
nn.SiLU(),
nn.Linear(hidden_size, hidden_size, bias=True)
)
self.mouse_qk_dim_list = mouse_qk_dim_list
self.heads_num = heads_num
if self.enable_mouse:
c = mouse_hidden_dim
self.mouse_mlp = nn.Sequential(
nn.Linear(mouse_dim_in * vae_time_compression_ratio * windows_size + img_hidden_size, c, bias=True),
nn.GELU(approximate="tanh"),
nn.Linear(c, c),
FP32LayerNorm(c, elementwise_affine=True),
)
head_dim = c // heads_num
self.t_qkv = ReplicatedLinear(c, c*3, bias=qkv_bias)
self.img_attn_q_norm = (
RMSNorm(head_dim, eps=1e-6)
if qk_norm
else nn.Identity()
)
self.img_attn_k_norm = (
RMSNorm(head_dim, eps=1e-6)
if qk_norm
else nn.Identity()
)
self.proj_mouse = ReplicatedLinear(c, img_hidden_size, bias=qkv_bias)
if self.enable_keyboard:
head_dim_key = keyboard_hidden_dim // heads_num
self.key_attn_q_norm = (
RMSNorm(head_dim_key, eps=1e-6)
if qk_norm
else nn.Identity()
)
self.key_attn_k_norm = (
RMSNorm(head_dim_key, eps=1e-6)
if qk_norm
else nn.Identity()
)
self.mouse_attn_q = ReplicatedLinear(img_hidden_size, keyboard_hidden_dim, bias=qkv_bias)
self.keyboard_attn_kv = ReplicatedLinear(hidden_size * windows_size * vae_time_compression_ratio, keyboard_hidden_dim * 2, bias=qkv_bias)
self.proj_keyboard = ReplicatedLinear(keyboard_hidden_dim, img_hidden_size, bias=qkv_bias)
self.mouse_attn_layer = LocalAttention(
num_heads=heads_num,
head_size=mouse_hidden_dim // heads_num,
causal=False,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)
) if self.enable_mouse else None
self.keyboard_attn_layer = LocalAttention(
num_heads=heads_num,
head_size=keyboard_hidden_dim // heads_num,
causal=False,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)
) if self.enable_keyboard else None
self.vae_time_compression_ratio = vae_time_compression_ratio
self.windows_size = windows_size
self.patch_size = patch_size
# Lazy initialization: freqs will be created on first forward pass
self._freqs_cos = None
self._freqs_sin = None
def patchify(self, x, patch_size):
"""
x : (N C T H W)
"""
pt, ph, pw = self.patch_size
t, h, w = x.shape[2] // pt, x.shape[3] // ph, x.shape[4] // pw
c = x.shape[1]
x = x.reshape(shape=(x.shape[0], c, t , pt, h , ph, w , pw))
x = torch.einsum("nctohpwq->nthwcopq", x)
x = x.reshape(shape=(x.shape[0], t*h*w, c*pt*ph*pw))
return x
def unpatchify(self, x, t, h, w, patch_size):
"""
x: (N, T, patch_size**2 * C)
imgs: (N, H, W, C)
"""
c = x.shape[2] // patch_size #self.unpatchify_channels
pt, ph, pw = self.patch_size
assert t * h * w == x.shape[1]
x = x.reshape(shape=(x.shape[0], t, h, w, c, pt, ph, pw))
x = torch.einsum("nthwcopq->nctohpwq", x)
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
return imgs
def get_rotary_pos_embed(self, video_length, height, width, head_dim, rope_dim_list = None, start_offset=0):
target_ndim = 3
ndim = 5 - 2
latents_size = [video_length+start_offset, height, width]
if isinstance(self.patch_size, int):
assert all(s % self.patch_size == 0 for s in latents_size), (
f"Latent size(last {ndim} dimensions) should be divisible by patch size({self.patch_size}), "
f"but got {latents_size}."
)
rope_sizes = [s // self.patch_size for s in latents_size]
elif isinstance(self.patch_size, list):
assert all(
s % self.patch_size[idx] == 0
for idx, s in enumerate(latents_size)
), (
f"Latent size(last {ndim} dimensions) should be divisible by patch size({self.patch_size}), "
f"but got {latents_size}."
)
rope_sizes = [
s // self.patch_size[idx] for idx, s in enumerate(latents_size)
]
if len(rope_sizes) != target_ndim:
rope_sizes = [1] * (target_ndim - len(rope_sizes)) + rope_sizes # time axis
if rope_dim_list is None:
rope_dim_list = [head_dim // target_ndim for _ in range(target_ndim)]
assert (
sum(rope_dim_list) == head_dim
), "sum(rope_dim_list) should equal to head_dim of attention layer"
# Use Matrix-Game wrapper for FastVideo's function
freqs_cos, freqs_sin = _get_nd_rotary_pos_embed_matrixgame(
rope_dim_list,
rope_sizes,
theta=self.rope_theta,
theta_rescale_factor=1,
)
return freqs_cos[-video_length*rope_sizes[1]*rope_sizes[2]//self.patch_size[0]:], freqs_sin[-video_length*rope_sizes[1]*rope_sizes[2]//self.patch_size[0]:]
def forward(self, x, tt, th, tw, mouse_condition=None, keyboard_condition=None, block_mask_mouse=None, block_mask_keyboard=None, is_causal=False, kv_cache_mouse=None, kv_cache_keyboard=None, start_frame=0, use_rope_keyboard=True, num_frame_per_block=3):
'''
hidden_states: B, tt*th*tw, C
mouse_condition: B, N_frames, C1
keyboard_condition: B, N_frames, C2
'''
assert use_rope_keyboard
B, N_frames, C = keyboard_condition.shape
assert tt*th*tw == x.shape[1]
assert ((N_frames - 1) + self.vae_time_compression_ratio) % self.vae_time_compression_ratio == 0
N_feats = int((N_frames - 1) / self.vae_time_compression_ratio) + 1
# Lazy initialization of freqs on first forward pass
if self._freqs_cos is None or self._freqs_sin is None:
self._freqs_cos, self._freqs_sin = self.get_rotary_pos_embed(
7500, self.patch_size[1], self.patch_size[2], 64,
self.mouse_qk_dim_list, start_offset=0
)
# Defined freqs_cis early so it's available for both mouse and keyboard
freqs_cis = (self._freqs_cos, self._freqs_sin)
assert (N_feats == tt and ((is_causal and kv_cache_mouse is None) or not is_causal)) or ((N_frames - 1) // self.vae_time_compression_ratio + 1 == start_frame + num_frame_per_block and is_causal)
if self.enable_mouse and mouse_condition is not None:
hidden_states = rearrange(x, "B (T S) C -> (B S) T C", T=tt, S=th*tw) # 65*272*480 -> 17*(272//16)*(480//16) -> 8670
B, N_frames, C = mouse_condition.shape
else:
hidden_states = x
# padding
pad_t = self.vae_time_compression_ratio * self.windows_size
if self.enable_mouse and mouse_condition is not None:
pad = mouse_condition[:, 0:1, :].expand(-1, pad_t, -1)
mouse_condition = torch.cat([pad, mouse_condition], dim=1)
if is_causal and kv_cache_mouse is not None:
mouse_condition = mouse_condition[:, self.vae_time_compression_ratio*(N_feats - num_frame_per_block - self.windows_size) + pad_t:, :]
group_mouse = [mouse_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(num_frame_per_block)]
else:
group_mouse = [mouse_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(N_feats)]
group_mouse = torch.stack(group_mouse, dim = 1)
S = th * tw
group_mouse = group_mouse.unsqueeze(-1).expand(B, num_frame_per_block, pad_t, C, S)
group_mouse = group_mouse.permute(0, 4, 1, 2, 3).reshape(B * S, num_frame_per_block, pad_t * C)
group_mouse = torch.cat([hidden_states, group_mouse], dim = -1)
group_mouse = self.mouse_mlp(group_mouse)
# qkv
mouse_qkv, _ = self.t_qkv(group_mouse)
q, k, v = rearrange(mouse_qkv, "B L (K H D) -> K B L H D", K=3, H=self.heads_num) # BHW F H C
q = self.img_attn_q_norm(q).to(v)
k = self.img_attn_k_norm(k).to(v)
# rope embd
# freqs_cis = (self.freqs_cos, self.freqs_sin)
q, k = _apply_rotary_emb_qk(q, k, freqs_cis[0], freqs_cis[1], start_offset=start_frame)
## TODO: adding cache here
if is_causal:
if kv_cache_mouse is None:
assert q.shape[0] == k.shape[0] and q.shape[0] % 880 == 0 # == 880, f"{q.shape[0]},{k.shape[0]}"
padded_length = math.ceil(q.shape[1] / 32) * 32 - q.shape[1]
padded_q = torch.cat(
[q,
torch.zeros([q.shape[0], padded_length, q.shape[2], q.shape[3]],
device=q.device, dtype=v.dtype)],
dim=1
)
padded_k = torch.cat(
[k, torch.zeros([k.shape[0], padded_length, k.shape[2], k.shape[3]],
device=k.device, dtype=v.dtype)],
dim=1
)
padded_v = torch.cat(
[v, torch.zeros([v.shape[0], padded_length, v.shape[2], v.shape[3]],
device=v.device, dtype=v.dtype)],
dim=1
)
attn = flex_attention(
query=padded_q.transpose(2, 1), # after: B, HW, F, C
key=padded_k.transpose(2, 1),
value=padded_v.transpose(2, 1),
block_mask=block_mask_mouse
)[:, :, :-padded_length].transpose(2, 1)
else:
current_start = start_frame
current_end = current_start + q.shape[1]
assert q.shape[1] == num_frame_per_block
sink_size = 0
max_attention_size = self.local_attn_size
sink_tokens = sink_size * 1
kv_cache_size = kv_cache_mouse["k"].shape[1]
num_new_tokens = q.shape[1]
if (current_end > kv_cache_mouse["global_end_index"].item()) and (
num_new_tokens + kv_cache_mouse["local_end_index"].item() > kv_cache_size):
num_evicted_tokens = num_new_tokens + kv_cache_mouse["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache_mouse["local_end_index"].item() - num_evicted_tokens - sink_tokens
kv_cache_mouse["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache_mouse["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
kv_cache_mouse["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache_mouse["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# Insert the new keys/values at the end
local_end_index = kv_cache_mouse["local_end_index"].item() + current_end - \
kv_cache_mouse["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
else:
local_end_index = kv_cache_mouse["local_end_index"].item() + current_end - kv_cache_mouse["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
kv_cache_mouse["k"][:, local_start_index:local_end_index] = k
kv_cache_mouse["v"][:, local_start_index:local_end_index] = v
attn = self.mouse_attn_layer(
q,
kv_cache_mouse["k"][:, max(0, local_end_index - max_attention_size):local_end_index],
kv_cache_mouse["v"][:, max(0, local_end_index - max_attention_size):local_end_index],
)
kv_cache_mouse["global_end_index"].fill_(current_end)
kv_cache_mouse["local_end_index"].fill_(local_end_index)
else:
attn = self.mouse_attn_layer(q, k, v)
# Compute cu_squlens and max_seqlen for flash attention
# qk norm
attn = rearrange(attn, '(b S) T h d -> b (T S) (h d)',b=B)
hidden_states = rearrange(x, "(B S) T C -> B (T S) C", B=B)
attn, _ = self.proj_mouse(attn)
hidden_states = hidden_states + attn
if self.enable_keyboard and keyboard_condition is not None:
pad = keyboard_condition[:, 0:1, :].expand(-1, pad_t, -1)
keyboard_condition = torch.cat([pad, keyboard_condition], dim=1)
if is_causal and kv_cache_keyboard is not None:
keyboard_condition = keyboard_condition[:, self.vae_time_compression_ratio*(N_feats - num_frame_per_block - self.windows_size) + pad_t:, :] # keyboard_condition[:, self.vae_time_compression_ratio*(start_frame - self.windows_size) + pad_t:start_frame * self.vae_time_compression_ratio + pad_t,:]
keyboard_condition = self.keyboard_embed(keyboard_condition)
group_keyboard = [keyboard_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(num_frame_per_block)]
else:
keyboard_condition = self.keyboard_embed(keyboard_condition)
group_keyboard = [keyboard_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(N_feats)]
group_keyboard = torch.stack(group_keyboard, dim = 1) # B F RW C
group_keyboard = group_keyboard.reshape(shape=(group_keyboard.shape[0],group_keyboard.shape[1],-1))
# apply cross attn
mouse_q, _ = self.mouse_attn_q(hidden_states)
keyboard_kv, _ = self.keyboard_attn_kv(group_keyboard)
B, L, HD = mouse_q.shape
D = HD // self.heads_num
q = mouse_q.view(B, L, self.heads_num, D)
B, L, KHD = keyboard_kv.shape
k, v = keyboard_kv.view(B, L, 2, self.heads_num, D).permute(2, 0, 1, 3, 4)
# Compute cu_squlens and max_seqlen for flash attention
# qk norm
q = self.key_attn_q_norm(q).to(v)
k = self.key_attn_k_norm(k).to(v)
S = th * tw
assert S == 880
# position embed
if use_rope_keyboard:
B, TS, H, D = q.shape
T_ = TS // S
q = q.view(B, T_, S, H, D).transpose(1, 2).reshape(B * S, T_, H, D)
q, k = _apply_rotary_emb_qk(q, k, freqs_cis[0], freqs_cis[1], start_offset=start_frame)
k1, k2, k3, k4 = k.shape
k = k.expand(S, k2, k3, k4)
v = v.expand(S, k2, k3, k4)
if is_causal:
if kv_cache_keyboard is None:
assert q.shape[0] == k.shape[0] and q.shape[0] % 880 == 0
padded_length = math.ceil(q.shape[1] / 32) * 32 - q.shape[1]
padded_q = torch.cat(
[q,
torch.zeros([q.shape[0], padded_length, q.shape[2], q.shape[3]],
device=q.device, dtype=v.dtype)],
dim=1
)
padded_k = torch.cat(
[k, torch.zeros([k.shape[0], padded_length, k.shape[2], k.shape[3]],
device=k.device, dtype=v.dtype)],
dim=1
)
padded_v = torch.cat(
[v, torch.zeros([v.shape[0], padded_length, v.shape[2], v.shape[3]],
device=v.device, dtype=v.dtype)],
dim=1
)
attn = flex_attention(
query=padded_q.transpose(2, 1), # after: B, HW, F, C
key=padded_k.transpose(2, 1),
value=padded_v.transpose(2, 1),
block_mask=block_mask_keyboard
)[:, :, :-padded_length].transpose(2, 1)
else:
current_start = start_frame
current_end = current_start + k.shape[1]
assert k.shape[1] == num_frame_per_block
sink_size = 0
max_attention_size = self.local_attn_size
sink_tokens = sink_size * 1
kv_cache_size = kv_cache_keyboard["k"].shape[1]
num_new_tokens = k.shape[1]
if (current_end > kv_cache_keyboard["global_end_index"].item()) and (
num_new_tokens + kv_cache_keyboard["local_end_index"].item() > kv_cache_size):
num_evicted_tokens = num_new_tokens + kv_cache_keyboard["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache_keyboard["local_end_index"].item() - num_evicted_tokens - sink_tokens
kv_cache_keyboard["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache_keyboard["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
kv_cache_keyboard["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache_keyboard["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# Insert the new keys/values at the end
local_end_index = kv_cache_keyboard["local_end_index"].item() + current_end - \
kv_cache_keyboard["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
else:
local_end_index = kv_cache_keyboard["local_end_index"].item() + current_end - kv_cache_keyboard["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
assert k.shape[0] == 880 # BS == 1 or the cache should not be saved/ load method should be modified
kv_cache_keyboard["k"][:, local_start_index:local_end_index] = k[:1]
kv_cache_keyboard["v"][:, local_start_index:local_end_index] = v[:1]
attn = self.keyboard_attn_layer(
q,
kv_cache_keyboard["k"][:, max(0, local_end_index - max_attention_size):local_end_index].repeat(S, 1, 1, 1),
kv_cache_keyboard["v"][:, max(0, local_end_index - max_attention_size):local_end_index].repeat(S, 1, 1, 1),
)
kv_cache_keyboard["global_end_index"].fill_(current_end)
kv_cache_keyboard["local_end_index"].fill_(local_end_index)
else:
attn = self.keyboard_attn_layer(q, k, v)
attn = rearrange(attn, '(B S) T H D -> B (T S) (H D)', S=S)
else:
if is_causal:
if kv_cache_keyboard is None:
padded_length = math.ceil(q.shape[1] / 32) * 32 - q.shape[1]
padded_q = torch.cat(
[q,
torch.zeros([q.shape[0], padded_length, q.shape[2], q.shape[3]],
device=q.device, dtype=v.dtype)],
dim=1
)
padded_k = torch.cat(
[k, torch.zeros([k.shape[0], padded_length, k.shape[2], k.shape[3]],
device=k.device, dtype=v.dtype)],
dim=1
)
padded_v = torch.cat(
[v, torch.zeros([v.shape[0], padded_length, v.shape[2], v.shape[3]],
device=v.device, dtype=v.dtype)],
dim=1
)
attn = flex_attention(
query=padded_q.transpose(2, 1), # after: B, HW, F, C
key=padded_k.transpose(2, 1),
value=padded_v.transpose(2, 1),
block_mask=block_mask_keyboard
)[:, :, :-padded_length].transpose(2, 1)
else:
current_start = start_frame
current_end = current_start + k.shape[1]
assert k.shape[1] == num_frame_per_block
sink_size = 0
max_attention_size = self.local_attn_size
sink_tokens = sink_size * 1
kv_cache_size = kv_cache_keyboard["k"].shape[1]
num_new_tokens = k.shape[1]
if (current_end > kv_cache_keyboard["global_end_index"].item()) and (
num_new_tokens + kv_cache_keyboard["local_end_index"].item() > kv_cache_size):
num_evicted_tokens = num_new_tokens + kv_cache_keyboard["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache_keyboard["local_end_index"].item() - num_evicted_tokens - sink_tokens
kv_cache_keyboard["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache_keyboard["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
kv_cache_keyboard["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache_keyboard["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# Insert the new keys/values at the end
local_end_index = kv_cache_keyboard["local_end_index"].item() + current_end - \
kv_cache_keyboard["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
else:
local_end_index = kv_cache_keyboard["local_end_index"].item() + current_end - kv_cache_keyboard["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
kv_cache_keyboard["k"][:, local_start_index:local_end_index] = k
kv_cache_keyboard["v"][:, local_start_index:local_end_index] = v
attn = self.keyboard_attn_layer(
q,
kv_cache_keyboard["k"][:, max(0, local_end_index - max_attention_size):local_end_index],
kv_cache_keyboard["v"][:, max(0, local_end_index - max_attention_size):local_end_index],
)
kv_cache_keyboard["global_end_index"].fill_(current_end)
kv_cache_keyboard["local_end_index"].fill_(local_end_index)
else:
attn = self.keyboard_attn_layer(q, k, v)
attn = rearrange(attn, 'B L H D -> B L (H D)')
attn, _ = self.proj_keyboard(attn)
hidden_states = hidden_states + attn
return hidden_states
File diff suppressed because it is too large Load Diff
+469
View File
@@ -0,0 +1,469 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any
import torch
import torch.nn as nn
from fastvideo.attention import DistributedAttention
from fastvideo.configs.models.dits.matrixgame import MatrixGameWanVideoConfig
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.layers.visual_embedding import PatchEmbed, TimestepEmbedder, ModulateProjection
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import (WanSelfAttention,
WanI2VCrossAttention,
WanT2VCrossAttention,
WanImageEmbedding)
from fastvideo.platforms import AttentionBackendEnum, current_platform
# Import ActionModule
from .action_module import ActionModule
logger = init_logger(__name__)
class MatrixGameTimeImageEmbedding(nn.Module):
def __init__(
self,
dim: int,
time_freq_dim: int,
image_embed_dim: int | None = None,
):
super().__init__()
self.time_embedder = TimestepEmbedder(
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
self.time_modulation = ModulateProjection(dim,
factor=6,
act_layer="silu")
self.image_embedder = None
if image_embed_dim is not None:
self.image_embedder = WanImageEmbedding(image_embed_dim, dim)
def forward(
self,
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor, # Kept for interface compatibility
encoder_hidden_states_image: torch.Tensor | None = None,
timestep_seq_len: int | None = None,
):
temb = self.time_embedder(timestep, timestep_seq_len)
timestep_proj = self.time_modulation(temb)
# MatrixGame does not use text embeddings, so we ignore encoder_hidden_states
# and return None for the text embedding part
if encoder_hidden_states_image is not None:
assert self.image_embedder is not None
encoder_hidden_states_image = self.image_embedder(
encoder_hidden_states_image)
return temb, timestep_proj, None, encoder_hidden_states_image
class MatrixGameCrossAttention(WanSelfAttention):
def forward(self, x, context, context_lens=None, crossattn_cache=None):
r"""
Args:
x(Tensor): Shape [B, L1, C]
context(Tensor): Shape [B, L2, C] - typically 257 image tokens
context_lens(Tensor): Shape [B]
crossattn_cache(dict): Optional cache for k/v during inference
"""
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)
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)
v = self.to_v(context)[0].view(b, -1, n, d)
crossattn_cache["k"] = k
crossattn_cache["v"] = v
else:
k = crossattn_cache["k"]
v = crossattn_cache["v"]
else:
k = self.norm_k(self.to_k(context)[0]).view(b, -1, n, d)
v = self.to_v(context)[0].view(b, -1, n, d)
# compute attention
x = self.attn(q, k, v)
# output
x = x.flatten(2)
x, _ = self.to_out(x)
return x
class MatrixGameTransformerBlock(nn.Module):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: int | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
prefix: str = "",
action_config: dict | None = None):
super().__init__()
action_config = action_config or {}
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
self.to_out = ReplicatedLinear(dim, dim, bias=True)
self.attn1 = DistributedAttention(
num_heads=num_heads,
head_size=dim // num_heads,
causal=False,
supported_attention_backends=supported_attention_backends,
prefix=f"{prefix}.attn1")
self.hidden_dim = dim
self.num_attention_heads = num_heads
dim_head = dim // num_heads
if qk_norm == "rms_norm":
self.norm_q = RMSNorm(dim_head, eps=eps)
self.norm_k = RMSNorm(dim_head, eps=eps)
elif qk_norm == "rms_norm_across_heads":
# LTX applies qk norm across all heads
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
else:
print("QK Norm type not supported")
raise Exception
assert cross_attn_norm is True
self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=True,
dtype=torch.float32,
compute_dtype=torch.float32)
# 2. Cross-attention
if added_kv_proj_dim is not None:
# I2V
self.attn2 = WanI2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
else:
# T2V
self.attn2 = WanT2VCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
# 2.1. Action Module Integration
self.use_action_module = len(action_config) > 0
if self.use_action_module:
self.action_model = ActionModule(**action_config)
else:
self.action_model = None
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
self.mlp_residual = ScaleResidual()
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
# Action Module specific args
grid_sizes: torch.Tensor | None = None,
mouse_cond: torch.Tensor | None = None,
keyboard_cond: torch.Tensor | None = None,
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
orig_dtype = hidden_states.dtype
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()
).chunk(6, dim=2)
# batch_size, seq_len, 1, inner_dim
shift_msa = shift_msa.squeeze(2)
scale_msa = scale_msa.squeeze(2)
gate_msa = gate_msa.squeeze(2)
c_shift_msa = c_shift_msa.squeeze(2)
c_scale_msa = c_scale_msa.squeeze(2)
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()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
# Apply rotary embeddings
cos, sin = freqs_cis
query, key = _apply_rotary_emb(query, cos, sin,
is_neox_style=False), _apply_rotary_emb(
key, cos, sin, is_neox_style=False)
attn_output, _ = self.attn1(query, key, value)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
context=encoder_hidden_states,
context_lens=None)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# ================= Action Module =================
if self.action_model is not None:
if mouse_cond is not None or keyboard_cond is not None:
# grid_sizes is expected to be [F, H, W]
# ActionModule implementation takes hidden_states directly
hidden_states = self.action_model(
hidden_states,
grid_sizes[0], grid_sizes[1], grid_sizes[2],
mouse_cond, keyboard_cond,
num_frame_per_block=grid_sizes[0],
)
# =================================================
# 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
_DEFAULT_MATRIXGAME_CONFIG = MatrixGameWanVideoConfig()
class MatrixGameWanModel(BaseDiT):
# Marker for action input support (Matrix-Game)
supports_action_input = True
_fsdp_shard_conditions = _DEFAULT_MATRIXGAME_CONFIG._fsdp_shard_conditions
_compile_conditions = _DEFAULT_MATRIXGAME_CONFIG._compile_conditions
_supported_attention_backends = _DEFAULT_MATRIXGAME_CONFIG._supported_attention_backends
param_names_mapping = _DEFAULT_MATRIXGAME_CONFIG.param_names_mapping
reverse_param_names_mapping = _DEFAULT_MATRIXGAME_CONFIG.reverse_param_names_mapping
lora_param_names_mapping = _DEFAULT_MATRIXGAME_CONFIG.lora_param_names_mapping
def __init__(self,
config: MatrixGameWanVideoConfig,
hf_config: dict[str, Any],
**kwargs) -> None:
super().__init__(config=config, hf_config=hf_config)
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.patch_size = config.patch_size
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
patch_size=config.patch_size,
flatten=False)
# 2. Condition embeddings
self.condition_embedder = MatrixGameTimeImageEmbedding(
dim=inner_dim,
time_freq_dim=config.freq_dim,
image_embed_dim=config.image_dim,
)
# 2.1. Get action config
self.action_config = getattr(config, 'action_config', {})
# 3. Transformer blocks
self.blocks = nn.ModuleList([
MatrixGameTransformerBlock(
inner_dim,
config.ffn_dim,
config.num_attention_heads,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
self._supported_attention_backends,
prefix=f"{getattr(config, 'prefix', 'Wan')}.blocks.{i}",
action_config=self.action_config)
for i in range(config.num_layers)
])
# 4. Output norm & projection
self.norm_out = LayerNormScaleShift(inner_dim,
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
self.gradient_checkpointing = False
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor
| list[torch.Tensor] | None = None,
# Action inputs
mouse_cond: torch.Tensor | None = None,
keyboard_cond: torch.Tensor | None = None,
**kwargs) -> torch.Tensor:
if encoder_hidden_states is not None and not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
else:
encoder_hidden_states_image = None
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# Get rotary embeddings
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sp_world_size(), post_patch_height,
post_patch_width),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
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
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
if timestep.dim() == 2:
timestep = timestep.flatten()
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))
if encoder_hidden_states is not None:
if isinstance(encoder_hidden_states, list):
encoder_hidden_states = encoder_hidden_states[0]
elif encoder_hidden_states.ndim == 2:
encoder_hidden_states = encoder_hidden_states.unsqueeze(0)
else:
# encoder_hidden_states is None (e.g. no text encoder)
# MatrixGame uses image-action cross-attn.
pass
if encoder_hidden_states_image is not None:
if encoder_hidden_states is not None:
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
else:
encoder_hidden_states = encoder_hidden_states_image
# This is [F, H, W] for the ActionModule
grid_sizes = torch.tensor([
post_patch_num_frames, post_patch_height, post_patch_width
],
device=hidden_states.device)
# Blocks
for block in self.blocks:
if torch.is_grad_enabled() and self.gradient_checkpointing:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states, timestep_proj,
freqs_cis,
grid_sizes=grid_sizes,
mouse_cond=mouse_cond,
keyboard_cond=keyboard_cond)
else:
hidden_states = block(
hidden_states,
encoder_hidden_states,
timestep_proj,
freqs_cis,
grid_sizes=grid_sizes,
mouse_cond=mouse_cond,
keyboard_cond=keyboard_cond)
# Output
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
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)
return output
+307
View File
@@ -0,0 +1,307 @@
from __future__ import annotations
import os
import random
# import cv2
import numpy as np
import torch
from diffusers.utils import export_to_video
from PIL import Image
from fastvideo.utils import logger
CAM_VALUE = 0.1
CAMERA_MAP = {
"i": [CAM_VALUE, 0], "k": [-CAM_VALUE, 0],
"j": [0, -CAM_VALUE], "l": [0, CAM_VALUE], "u": [0, 0]
}
KEYBOARD_MAP_4 = { # base_distilled_model (universal): W/S/A/D
"w": [1, 0, 0, 0], "s": [0, 1, 0, 0],
"a": [0, 0, 1, 0], "d": [0, 0, 0, 1], "q": [0, 0, 0, 0]
}
KEYBOARD_MAP_2 = { # gta_distilled_model: W/S only (steering via mouse)
"w": [1, 0], "s": [0, 1], "q": [0, 0]
}
KEYBOARD_MAP_7 = { # templerun_distilled_model: still/w/s/left/right/a/d
"q": [1, 0, 0, 0, 0, 0, 0], # still
"w": [0, 1, 0, 0, 0, 0, 0], # forward
"s": [0, 0, 1, 0, 0, 0, 0], # back
"j": [0, 0, 0, 1, 0, 0, 0], # left (swipe)
"l": [0, 0, 0, 0, 1, 0, 0], # right (swipe)
"a": [0, 0, 0, 0, 0, 1, 0], # a
"d": [0, 0, 0, 0, 0, 0, 1], # d
}
KEYBOARD_MAP = KEYBOARD_MAP_4 # Default for backward compatibility
def load_initial_image(image_path: str = None) -> Image.Image:
if image_path and os.path.exists(image_path):
return Image.open(image_path).convert("RGB")
logger.warning("No image provided, creating placeholder...")
return Image.new("RGB", (640, 352), (128, 128, 128))
def create_action_presets(num_frames: int, keyboard_dim: int = 4, seed: int = None):
if keyboard_dim not in (2, 4, 7):
raise ValueError(f"keyboard_dim must be 2, 4, or 7, got {keyboard_dim}")
if num_frames % 4 != 1:
raise ValueError("Matrix-Game conditioning expects num_frames to be 4k+1.")
# Set seed for reproducibility if provided
if seed is not None:
random.seed(seed)
num_samples_per_action = 4
# Define actions based on keyboard_dim
if keyboard_dim == 4:
# Universal model: W, S, A, D
actions_single_action = ["forward", "left", "right"]
actions_double_action = ["forward_left", "forward_right"]
actions_single_camera = ["camera_l", "camera_r"]
keyboard_idx = {"forward": 0, "back": 1, "left": 2, "right": 3}
elif keyboard_dim == 2:
# GTA model: W, S only (steering via mouse)
actions_single_action = ["forward", "back"]
actions_double_action = []
actions_single_camera = ["camera_l", "camera_r"]
keyboard_idx = {"forward": 0, "back": 1}
else: # keyboard_dim == 7
# Temple Run model: still, w, s, left, right, a, d (no mouse)
actions_single_action = ["forward", "back", "left", "right"]
actions_double_action = []
actions_single_camera = [] # No mouse for Temple Run
keyboard_idx = {"still": 0, "forward": 1, "back": 2, "left": 3, "right": 4, "a": 5, "d": 6}
actions_to_test = (
actions_double_action * 5 + actions_single_camera * 5 + actions_single_action * 5
)
for action in (actions_single_action + actions_double_action):
for camera in actions_single_camera:
actions_to_test.append(f"{action}_{camera}")
# Ensure we have at least some actions
if not actions_to_test:
actions_to_test = actions_single_action * 5
base_action = actions_single_action + actions_single_camera
cam_value = 0.1
camera_value_map = {
"camera_up": [cam_value, 0],
"camera_down": [-cam_value, 0],
"camera_l": [0, -cam_value],
"camera_r": [0, cam_value],
"camera_ur": [cam_value, cam_value],
"camera_ul": [cam_value, -cam_value],
"camera_dr": [-cam_value, cam_value],
"camera_dl": [-cam_value, -cam_value],
}
data = []
for action_name in actions_to_test:
keyboard_condition = torch.zeros((num_samples_per_action, keyboard_dim))
mouse_condition = torch.zeros((num_samples_per_action, 2))
for sub_act in base_action:
if sub_act not in action_name:
continue
if sub_act in camera_value_map:
mouse_condition = torch.tensor(
[camera_value_map[sub_act] for _ in range(num_samples_per_action)],
dtype=mouse_condition.dtype,
)
elif sub_act in keyboard_idx:
keyboard_condition[:, keyboard_idx[sub_act]] = 1
data.append({
"keyboard_condition": keyboard_condition,
"mouse_condition": mouse_condition,
})
keyboard_condition = torch.zeros((num_frames, keyboard_dim))
mouse_condition = torch.zeros((num_frames, 2))
current_frame = 0
selections = [12]
while current_frame < num_frames:
rd_frame = selections[random.randint(0, len(selections) - 1)]
entry = data[random.randint(0, len(data) - 1)]
key_seq = entry["keyboard_condition"]
mouse_seq = entry["mouse_condition"]
if current_frame == 0:
keyboard_condition[:1] = key_seq[:1]
mouse_condition[:1] = mouse_seq[:1]
current_frame = 1
else:
rd_frame = min(rd_frame, num_frames - current_frame)
repeat_time = rd_frame // 4
keyboard_condition[current_frame:current_frame + rd_frame] = key_seq.repeat(repeat_time, 1)
mouse_condition[current_frame:current_frame + rd_frame] = mouse_seq.repeat(repeat_time, 1)
current_frame += rd_frame
return {"keyboard": keyboard_condition, "mouse": mouse_condition}
def parse_config(config, mode="universal"):
assert mode in ['universal', 'gta_drive', 'templerun']
key_data = {}
mouse_data = {}
if mode != 'templerun':
key, mouse = config
else:
key = config
for i in range(len(key)):
if mode == 'templerun':
still, w, s, left, right, a, d = key[i]
elif mode == 'universal':
w, s, a, d = key[i]
else:
w, s, a, d = key[i][0], key[i][1], mouse[i][1] < 0, mouse[i][1] > 0
if mode == 'universal':
mouse_y, mouse_x = mouse[i]
mouse_y = -1 * mouse_y
key_data[i] = {"W": bool(w), "A": bool(a), "S": bool(s), "D": bool(d)}
if mode == 'templerun':
key_data[i].update({"left": bool(left), "right": bool(right)})
if mode == 'universal':
if i == 0:
mouse_data[i] = (320, 352 // 2)
else:
global_scale_factor = 0.1
mouse_scale_x = 15 * global_scale_factor
mouse_scale_y = 15 * 4 * global_scale_factor
mouse_data[i] = (
mouse_data[i - 1][0] + mouse_x * mouse_scale_x,
mouse_data[i - 1][1] + mouse_y * mouse_scale_y,
)
return key_data, mouse_data
# NOTE: drawing functions are commented out to avoid cv2/libGL dependency.
#
# def draw_rounded_rectangle(image, top_left, bottom_right, color, radius=10, alpha=0.5):
# overlay = image.copy()
# x1, y1 = top_left
# x2, y2 = bottom_right
#
# cv2.rectangle(overlay, (x1 + radius, y1), (x2 - radius, y2), color, -1)
# cv2.rectangle(overlay, (x1, y1 + radius), (x2, y2 - radius), color, -1)
# cv2.ellipse(overlay, (x1 + radius, y1 + radius), (radius, radius), 180, 0, 90, color, -1)
# cv2.ellipse(overlay, (x2 - radius, y1 + radius), (radius, radius), 270, 0, 90, color, -1)
# cv2.ellipse(overlay, (x1 + radius, y2 - radius), (radius, radius), 90, 0, 90, color, -1)
# cv2.ellipse(overlay, (x2 - radius, y2 - radius), (radius, radius), 0, 0, 90, color, -1)
# cv2.addWeighted(overlay, alpha, image, 1 - alpha, 0, image)
#
#
# def draw_keys_on_frame(frame, keys, key_size=(80, 50), spacing=20, bottom_margin=30, mode='universal'):
# h, w, _ = frame.shape
# horison_shift = 90
# vertical_shift = -20
# horizon_shift_all = 50
# key_positions = {
# "W": (w // 2 - key_size[0] // 2 - horison_shift - horizon_shift_all,
# h - bottom_margin - key_size[1] * 2 + vertical_shift - 20),
# "A": (w // 2 - key_size[0] * 2 + 5 - horison_shift - horizon_shift_all,
# h - bottom_margin - key_size[1] + vertical_shift),
# "S": (w // 2 - key_size[0] // 2 - horison_shift - horizon_shift_all,
# h - bottom_margin - key_size[1] + vertical_shift),
# "D": (w // 2 + key_size[0] - 5 - horison_shift - horizon_shift_all,
# h - bottom_margin - key_size[1] + vertical_shift),
# }
# key_icon = {"W": "W", "A": "A", "S": "S", "D": "D", "left": "left", "right": "right"}
# if mode == 'templerun':
# key_positions.update({
# "left": (w // 2 + key_size[0] * 2 + spacing * 2 - horison_shift - horizon_shift_all,
# h - bottom_margin - key_size[1] + vertical_shift),
# "right": (w // 2 + key_size[0] * 3 + spacing * 7 - horison_shift - horizon_shift_all,
# h - bottom_margin - key_size[1] + vertical_shift)
# })
#
# for key, (x, y) in key_positions.items():
# is_pressed = keys.get(key, False)
# top_left = (x, y)
# if key in ["left", "right"]:
# bottom_right = (x + key_size[0] + 40, y + key_size[1])
# else:
# bottom_right = (x + key_size[0], y + key_size[1])
#
# color = (0, 255, 0) if is_pressed else (200, 200, 200)
# alpha = 0.8 if is_pressed else 0.5
# draw_rounded_rectangle(frame, top_left, bottom_right, color, radius=10, alpha=alpha)
#
# text_size = cv2.getTextSize(key, cv2.FONT_HERSHEY_SIMPLEX, 0.8, 2)[0]
# if key in ["left", "right"]:
# text_x = x + (key_size[0] + 40 - text_size[0]) // 2
# else:
# text_x = x + (key_size[0] - text_size[0]) // 2
# text_y = y + (key_size[1] + text_size[1]) // 2
# cv2.putText(frame, key_icon[key], (text_x, text_y), cv2.FONT_HERSHEY_SIMPLEX, 0.8, (0, 0, 0), 2)
#
#
# def overlay_icon(frame, icon, position, scale=1.0, rotation=0):
# x, y = position
# h, w, _ = icon.shape
#
# scaled_width = int(w * scale)
# scaled_height = int(h * scale)
# icon_resized = cv2.resize(icon, (scaled_width, scaled_height), interpolation=cv2.INTER_AREA)
#
# center = (scaled_width // 2, scaled_height // 2)
# rotation_matrix = cv2.getRotationMatrix2D(center, rotation, 1.0)
# icon_rotated = cv2.warpAffine(
# icon_resized, rotation_matrix, (scaled_width, scaled_height),
# flags=cv2.INTER_LINEAR, borderMode=cv2.BORDER_CONSTANT, borderValue=(0, 0, 0, 0)
# )
#
# h, w, _ = icon_rotated.shape
# frame_h, frame_w, _ = frame.shape
#
# top_left_x = max(0, int(x - w // 2))
# top_left_y = max(0, int(y - h // 2))
# bottom_right_x = min(frame_w, int(x + w // 2))
# bottom_right_y = min(frame_h, int(y + h // 2))
#
# icon_x_start = max(0, int(-x + w // 2))
# icon_y_start = max(0, int(-y + h // 2))
# icon_x_end = icon_x_start + (bottom_right_x - top_left_x)
# icon_y_end = icon_y_start + (bottom_right_y - top_left_y)
#
# icon_region = icon_rotated[icon_y_start:icon_y_end, icon_x_start:icon_x_end]
# alpha = icon_region[:, :, 3] / 255.0
# icon_rgb = icon_region[:, :, :3]
#
# frame_region = frame[top_left_y:bottom_right_y, top_left_x:bottom_right_x]
# for c in range(3):
# frame_region[:, :, c] = (1 - alpha) * frame_region[:, :, c] + alpha * icon_rgb[:, :, c]
# frame[top_left_y:bottom_right_y, top_left_x:bottom_right_x] = frame_region
#
#
# def process_video(input_video, output_video, config, mouse_icon_path,
# mouse_scale=1.0, mouse_rotation=0, process_icon=True, mode='universal'):
# key_data, mouse_data = parse_config(config, mode=mode)
# fps = 12
#
# mouse_icon = cv2.imread(mouse_icon_path, cv2.IMREAD_UNCHANGED)
#
# out_video = []
# for frame_idx, frame in enumerate(input_video):
# frame = np.ascontiguousarray(frame)
# if process_icon:
# keys = key_data.get(frame_idx, {"W": False, "A": False, "S": False, "D": False, "left": False, "right": False})
# draw_keys_on_frame(frame, keys, key_size=(50, 50), spacing=10, bottom_margin=20, mode=mode)
# if mode == 'universal':
# frame_width = frame.shape[1]
# frame_height = frame.shape[0]
# mouse_position = mouse_data.get(frame_idx, (frame_width // 2, frame_height // 2))
# overlay_icon(frame, mouse_icon, mouse_position, scale=mouse_scale, rotation=mouse_rotation)
# out_video.append(frame / 255)
#
# export_to_video(out_video, output_video, fps=fps)
# logger.info(f"Video saved to {output_video}")
+53 -24
View File
@@ -12,7 +12,10 @@ from fastvideo.attention import (DistributedAttention, DistributedAttention_VSA,
LocalAttention)
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.configs.sample.wan import WanTeaCacheParams
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather,
sequence_model_parallel_all_gather_with_unpad,
sequence_model_parallel_shard)
from fastvideo.forward_context import get_forward_context
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
@@ -21,14 +24,16 @@ from fastvideo.layers.linear import ReplicatedLinear
# from torch.nn import RMSNorm
# TODO: RMSNorm ....
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
from fastvideo.layers.visual_embedding import (ModulateProjection, PatchEmbed,
TimestepEmbedder)
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import CachableDiT
from fastvideo.platforms import AttentionBackendEnum, current_platform
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.distributed.utils import create_attention_mask_for_padding
logger = init_logger(__name__)
@@ -314,6 +319,7 @@ class WanTransformerBlock(nn.Module):
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
attention_mask: torch.Tensor | None = None,
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
@@ -356,13 +362,7 @@ class WanTransformerBlock(nn.Module):
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
# Apply rotary embeddings
cos, sin = freqs_cis
query, key = _apply_rotary_emb(query, cos, sin,
is_neox_style=False), _apply_rotary_emb(
key, cos, sin, is_neox_style=False)
attn_output, _ = self.attn1(query, key, value)
attn_output, _ = self.attn1(query, key, value, freqs_cis=freqs_cis, attention_mask=attention_mask)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
@@ -474,6 +474,7 @@ class WanTransformerBlock_VSA(nn.Module):
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
attention_mask: torch.Tensor | None = None,
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
@@ -504,16 +505,12 @@ class WanTransformerBlock_VSA(nn.Module):
gate_compress = gate_compress.squeeze(1).unflatten(
2, (self.num_attention_heads, -1))
# Apply rotary embeddings
cos, sin = freqs_cis
query, key = _apply_rotary_emb(query, cos, sin,
is_neox_style=False), _apply_rotary_emb(
key, cos, sin, is_neox_style=False)
attn_output, _ = self.attn1(query,
key,
value,
gate_compress=gate_compress)
freqs_cis = freqs_cis,
gate_compress=gate_compress,
attention_mask=attention_mask)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
@@ -563,6 +560,8 @@ class WanTransformer3DModel(CachableDiT):
self.patch_size = config.patch_size
self.text_len = config.text_len
assert config.num_attention_heads % get_sp_world_size() == 0, f"The number of attention heads ({config.num_attention_heads}) must be divisible by the sequence parallel size ({get_sp_world_size()})"
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
@@ -606,6 +605,7 @@ class WanTransformer3DModel(CachableDiT):
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
self.gradient_checkpointing = False
self._logged_attention_mask = False
# For type checking
self.previous_e0_even = None
@@ -650,21 +650,43 @@ class WanTransformer3DModel(CachableDiT):
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sp_world_size(), post_patch_height,
(post_patch_num_frames, post_patch_height,
post_patch_width),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
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.to(hidden_states.device).float(),
freqs_sin.to(hidden_states.device).float())
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# Shard with padding support - returns (sharded_tensor, original_seq_len)
hidden_states, original_seq_len = sequence_model_parallel_shard(hidden_states, dim=1)
# Create attention mask for padded tokens if padding was applied
current_seq_len = hidden_states.shape[1]
sp_world_size = get_sp_world_size()
padded_seq_len = current_seq_len * sp_world_size
if padded_seq_len > original_seq_len:
if not self._logged_attention_mask:
logger.info(f"Padding applied, original seq len: {original_seq_len}, padded seq len: {padded_seq_len}")
self._logged_attention_mask = True
attention_mask = create_attention_mask_for_padding(
seq_len=original_seq_len,
padded_seq_len=padded_seq_len,
batch_size=batch_size,
device=hidden_states.device,
)
else:
if not self._logged_attention_mask:
logger.info(f"Padding not applied")
self._logged_attention_mask = True
attention_mask = None
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
@@ -698,6 +720,7 @@ class WanTransformer3DModel(CachableDiT):
timestep_proj=timestep_proj, temb=temb)
if should_skip_forward:
print("skipping forward, cached")
hidden_states = self.retrieve_cached_states(hidden_states)
else:
# if teacache is enabled, we need to cache the original hidden states
@@ -708,12 +731,13 @@ class WanTransformer3DModel(CachableDiT):
for block in self.blocks:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis)
timestep_proj, freqs_cis, attention_mask)
else:
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis)
timestep_proj, freqs_cis, attention_mask)
# if teacache is enabled, we need to cache the original hidden states
if enable_teacache:
self.maybe_cache_states(hidden_states, original_hidden_states)
# 5. Output norm, projection & unpatchify
@@ -726,7 +750,12 @@ class WanTransformer3DModel(CachableDiT):
# batch_size, inner_dim
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
hidden_states = self.norm_out(hidden_states, shift, scale)
# Gather and unpad in one operation
hidden_states = sequence_model_parallel_all_gather_with_unpad(
hidden_states, original_seq_len, dim=1)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
+387
View File
@@ -0,0 +1,387 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from transformers: https://github.com/huggingface/transformers/blob/main/src/transformers/models/qwen2_5_vl/modeling_qwen2_5_vl.py
import math
from typing import Any, Optional, Tuple, Union, List, Callable, Iterable
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.configs.models.encoders import BaseEncoderOutput, Qwen2_5_VLConfig
from fastvideo.distributed import get_tp_rank, get_tp_world_size
from fastvideo.layers.activation import get_act_fn, SiluAndMul
from fastvideo.layers.layernorm import RMSNorm
from fastvideo.layers.linear import MergedColumnParallelLinear, QKVParallelLinear, RowParallelLinear
from fastvideo.layers.quantization import QuantizationConfig
from fastvideo.layers.vocab_parallel_embedding import VocabParallelEmbedding
from fastvideo.models.encoders.base import TextEncoder
from fastvideo.models.loader.weight_utils import default_weight_loader
from fastvideo.models.mask_utils import sdpa_mask
from fastvideo.logger import init_logger
logger = init_logger(__name__)
def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
"""
This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch,
num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim)
"""
batch, num_key_value_heads, slen, head_dim = hidden_states.shape
if n_rep == 1:
return hidden_states
hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim)
return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim)
def sdpa_attention_forward(
module: torch.nn.Module,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attention_mask: Optional[torch.Tensor],
dropout: float = 0.0,
scaling: Optional[float] = None,
is_causal: Optional[bool] = None,
**kwargs,
) -> tuple[torch.Tensor, None]:
if kwargs.get("output_attentions", False) or kwargs.get("head_mask") is not None:
logger.warning_once(
"`sdpa` attention does not support `output_attentions=True` or `head_mask`."
" Please set your attention to `eager` if you want any of these features."
)
if hasattr(module, "num_key_value_groups"):
key = repeat_kv(key, module.num_key_value_groups)
value = repeat_kv(value, module.num_key_value_groups)
if attention_mask is not None and attention_mask.ndim == 4:
attention_mask = attention_mask[:, :, :, : key.shape[-2]]
# If attention_mask is not None, convert it to boolean type
if attention_mask is not None and attention_mask.dtype != torch.bool:
attention_mask = attention_mask.bool()
attn_output = torch.nn.functional.scaled_dot_product_attention(
query,
key,
value,
attn_mask=attention_mask,
dropout_p=dropout,
scale=scaling,
is_causal=is_causal,
)
attn_output = attn_output.transpose(1, 2).contiguous()
return attn_output, None
def rotate_half(x):
"""Rotates half the hidden dims of the input."""
x1 = x[..., : x.shape[-1] // 2]
x2 = x[..., x.shape[-1] // 2 :]
return torch.cat((-x2, x1), dim=-1)
def apply_multimodal_rotary_pos_emb(q, k, cos, sin, mrope_section, unsqueeze_dim=1):
mrope_section = [s * 2 for s in mrope_section]
cos = torch.cat([m[i % 3] for i, m in enumerate(cos.split(mrope_section, dim=-1))], dim=-1).unsqueeze(
unsqueeze_dim
)
sin = torch.cat([m[i % 3] for i, m in enumerate(sin.split(mrope_section, dim=-1))], dim=-1).unsqueeze(
unsqueeze_dim
)
q_embed = (q * cos) + (rotate_half(q) * sin)
k_embed = (k * cos) + (rotate_half(k) * sin)
return q_embed, k_embed
class Qwen2_5_VLRotaryEmbedding(nn.Module):
def __init__(self, config: Qwen2_5_VLConfig, device=None):
super().__init__()
self.max_seq_len_cached = config.max_position_embeddings
self.original_max_seq_len = config.max_position_embeddings
self.config = config
self.rope_type = config.rope_scaling.get("rope_type", "default")
self.base = config.rope_theta
# Simplified initialization
head_dim = config.hidden_size // config.num_attention_heads
dim = head_dim
self.attention_scaling = 1.0
inv_freq = 1.0 / (
self.base ** (torch.arange(0, dim, 2, dtype=torch.int64).float().to(device) / dim)
)
self.register_buffer("inv_freq", inv_freq, persistent=False)
self.original_inv_freq = inv_freq
def forward(self, x, position_ids):
# In contrast to other models, Qwen2_5_VL has different position ids for the grids
# So we expand the inv_freq to shape (3, ...)
inv_freq_expanded = self.inv_freq[None, None, :, None].float().expand(3, position_ids.shape[1], -1, 1)
position_ids_expanded = position_ids[:, :, None, :].float() # shape (3, bs, 1, positions)
device_type = x.device.type if isinstance(x.device.type, str) and x.device.type != "mps" else "cpu"
with torch.autocast(device_type=device_type, enabled=False): # Force float32
freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(2, 3)
emb = torch.cat((freqs, freqs), dim=-1)
cos = emb.cos() * self.attention_scaling
sin = emb.sin() * self.attention_scaling
return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype)
class Qwen2_5_VLMLP(nn.Module):
def __init__(self, config: Qwen2_5_VLConfig, quant_config: QuantizationConfig | None = None, prefix: str = ""):
super().__init__()
self.gate_up_proj = MergedColumnParallelLinear(
input_size=config.hidden_size,
output_sizes=[config.intermediate_size] * 2,
bias=False,
quant_config=quant_config,
prefix=f"{prefix}.gate_up_proj",
)
self.down_proj = RowParallelLinear(
input_size=config.intermediate_size,
output_size=config.hidden_size,
bias=False,
quant_config=quant_config,
prefix=f"{prefix}.down_proj",
)
self.act_fn = SiluAndMul()
def forward(self, x):
x, _ = self.gate_up_proj(x)
x = self.act_fn(x)
x, _ = self.down_proj(x)
return x
class Qwen2_5_VLAttention(nn.Module):
def __init__(self, config: Qwen2_5_VLConfig, layer_idx: int, quant_config: QuantizationConfig | None = None, prefix: str = ""):
super().__init__()
self.config = config
self.layer_idx = layer_idx
self.hidden_size = config.hidden_size
self.num_heads = config.num_attention_heads
self.head_dim = self.hidden_size // self.num_heads
self.num_key_value_heads = config.num_key_value_heads
self.num_key_value_groups = self.num_heads // self.num_key_value_heads
tp_size = get_tp_world_size()
self.total_num_heads = self.num_heads
assert self.total_num_heads % tp_size == 0
self.num_heads = self.total_num_heads // tp_size
self.total_num_kv_heads = self.num_key_value_heads
if self.total_num_kv_heads >= tp_size:
assert self.total_num_kv_heads % tp_size == 0
self.num_kv_heads = self.total_num_kv_heads // tp_size
else:
assert tp_size % self.total_num_kv_heads == 0
self.num_kv_heads = 1
self.num_kv_heads = max(1, self.total_num_kv_heads // tp_size)
self.q_size = self.num_heads * self.head_dim
self.kv_size = self.num_kv_heads * self.head_dim
self.scaling = self.head_dim**-0.5
self.qkv_proj = QKVParallelLinear(
hidden_size=self.hidden_size,
head_size=self.head_dim,
total_num_heads=self.total_num_heads,
total_num_kv_heads=self.total_num_kv_heads,
bias=True,
quant_config=quant_config,
prefix=f"{prefix}.qkv_proj",
)
self.o_proj = RowParallelLinear(
input_size=self.total_num_heads * self.head_dim,
output_size=self.hidden_size,
bias=False,
quant_config=quant_config,
prefix=f"{prefix}.o_proj",
)
self.layer_type = config.layer_types[layer_idx] if config.layer_types else "full_attention"
self.sliding_window = config.sliding_window if self.layer_type == "sliding_attention" else None
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.LongTensor] = None,
position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None,
output_attentions: bool = False,
) -> Tuple[torch.Tensor, Optional[torch.Tensor]]:
bsz, q_len, _ = hidden_states.size()
qkv, _ = self.qkv_proj(hidden_states)
query_states, key_states, value_states = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1)
query_states = query_states.view(bsz, q_len, self.num_heads, self.head_dim).transpose(1, 2)
key_states = key_states.view(bsz, q_len, self.num_kv_heads, self.head_dim).transpose(1, 2)
value_states = value_states.view(bsz, q_len, self.num_kv_heads, self.head_dim).transpose(1, 2)
cos, sin = position_embeddings
query_states, key_states = apply_multimodal_rotary_pos_emb(
query_states, key_states, cos, sin, self.config.rope_scaling["mrope_section"]
)
attn_output = sdpa_attention_forward(self, query_states, key_states, value_states, attention_mask, dropout=self.config.attention_dropout, scaling=self.scaling, is_causal=False)[0].reshape(bsz, q_len, -1)
attn_output, _ = self.o_proj(attn_output)
return attn_output
class Qwen2_5_VLDecoderLayer(nn.Module):
def __init__(self, config: Qwen2_5_VLConfig, layer_idx: int, quant_config: QuantizationConfig | None = None, prefix: str = ""):
super().__init__()
self.self_attn = Qwen2_5_VLAttention(config, layer_idx, quant_config=quant_config, prefix=f"{prefix}.self_attn")
self.mlp = Qwen2_5_VLMLP(config, quant_config=quant_config, prefix=f"{prefix}.mlp")
self.input_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.post_attention_layernorm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
def forward(
self,
hidden_states: torch.Tensor,
attention_mask: Optional[torch.Tensor] = None,
position_ids: Optional[torch.LongTensor] = None,
position_embeddings: Optional[tuple[torch.Tensor, torch.Tensor]] = None,
output_attentions: Optional[bool] = False,
) -> Tuple[torch.FloatTensor, Optional[Tuple[torch.FloatTensor, torch.FloatTensor]]]:
residual = hidden_states
hidden_states = self.input_layernorm(hidden_states)
hidden_states = self.self_attn(
hidden_states=hidden_states,
attention_mask=attention_mask,
position_ids=position_ids,
position_embeddings=position_embeddings,
output_attentions=output_attentions,
)
hidden_states = residual + hidden_states
residual = hidden_states
hidden_states = self.post_attention_layernorm(hidden_states)
hidden_states = self.mlp(hidden_states)
hidden_states = residual + hidden_states
outputs = (hidden_states,)
return outputs
class Qwen2_5_VLTextModel(TextEncoder):
def __init__(self, config: Qwen2_5_VLConfig):
super().__init__(config)
quant_config = None
self.embed_tokens = VocabParallelEmbedding(
config.vocab_size,
config.hidden_size,
org_num_embeddings=config.vocab_size
)
self.layers = nn.ModuleList([
Qwen2_5_VLDecoderLayer(config, layer_idx, quant_config=quant_config, prefix=f"{config.prefix}.layers.{layer_idx}")
for layer_idx in range(config.num_hidden_layers)
])
self.norm = RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
self.rotary_emb = Qwen2_5_VLRotaryEmbedding(config)
def get_input_embeddings(self, input_ids: torch.Tensor) -> torch.Tensor:
return self.embed_tokens(input_ids)
def forward(
self,
input_ids: Optional[torch.LongTensor] = None,
position_ids: Optional[torch.LongTensor] = None,
attention_mask: Optional[torch.Tensor] = None,
inputs_embeds: Optional[torch.FloatTensor] = None,
output_hidden_states: Optional[bool] = None,
**kwargs,
) -> BaseEncoderOutput:
output_hidden_states = (
output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states
)
if inputs_embeds is None:
inputs_embeds = self.get_input_embeddings(input_ids)
hidden_states = inputs_embeds
if position_ids is None:
seq_length = hidden_states.shape[1]
cache_position = torch.arange(seq_length, device=hidden_states.device)
position_ids = cache_position.view(1, 1, -1).expand(3, hidden_states.shape[0], -1)
mask_kwargs = {
"batch_size": hidden_states.shape[0],
"cache_position": cache_position,
"kv_length": attention_mask.shape[-1],
"kv_offset": 0,
"attention_mask": attention_mask
}
position_embeddings = self.rotary_emb(hidden_states, position_ids)
all_hidden_states = () if output_hidden_states else None
for decoder_layer in self.layers:
if output_hidden_states:
all_hidden_states += (hidden_states,)
layer_outputs = decoder_layer(
hidden_states,
attention_mask=sdpa_mask(**mask_kwargs),
position_ids=position_ids,
position_embeddings=position_embeddings,
output_attentions=False,
)
hidden_states = layer_outputs[0]
hidden_states = self.norm(hidden_states)
if output_hidden_states:
all_hidden_states += (hidden_states,)
return BaseEncoderOutput(
last_hidden_state=hidden_states,
hidden_states=all_hidden_states,
)
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
params_dict = dict(self.named_parameters())
loaded_params: set[str] = set()
for name, loaded_weight in weights:
if "rotary_emb.inv_freq" in name:
continue
for param_name, weight_name, shard_id in self.config.stacked_params_mapping:
if weight_name not in name:
continue
name = name.replace(weight_name, param_name)
# Skip loading extra bias for GPTQ models.
# if name.endswith(".bias") and name not in params_dict:
# continue
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = param.weight_loader
weight_loader(param, loaded_weight, shard_id)
break
else:
# Skip loading extra bias for GPTQ models.
# if name.endswith(".bias") and name not in params_dict:
# continue
if name not in params_dict:
continue
param = params_dict[name]
weight_loader = getattr(param, "weight_loader", default_weight_loader)
weight_loader(param, loaded_weight)
loaded_params.add(name)
return loaded_params
+2 -1
View File
@@ -19,7 +19,8 @@ import torch
import torch.nn as nn
import torch.nn.functional as F
from einops import rearrange
from transformers.modeling_utils import PretrainedConfig, PreTrainedModel
from transformers import PretrainedConfig
from transformers.modeling_utils import PreTrainedModel
from fastvideo.models.dits.stepvideo import StepVideoRMSNorm
+3 -7
View File
@@ -412,11 +412,9 @@ class VAELoader(ComponentLoader):
# Find all safetensors files
safetensors_list = glob.glob(
os.path.join(str(model_path), "*.safetensors"))
# TODO(PY)
assert len(
safetensors_list
) == 1, f"Found {len(safetensors_list)} safetensors files in {model_path}"
loaded = safetensors_load_file(safetensors_list[0])
loaded = {}
for sf_file in safetensors_list:
loaded.update(safetensors_load_file(sf_file))
vae.load_state_dict(
loaded, strict=False) # We might only load encoder or decoder
@@ -478,8 +476,6 @@ class TransformerLoader(ComponentLoader):
fastvideo_args.pipeline_config.dit_precision]
# Load the model using FSDP loader
logger.info("Loading model from %s, default_dtype: %s", cls_name,
default_dtype)
assert fastvideo_args.hsdp_shard_dim is not None
model = maybe_load_fsdp_model(
model_cls=model_cls,
+201
View File
@@ -0,0 +1,201 @@
import torch
from typing import Callable, Optional
def and_masks(*mask_functions: Callable) -> Callable:
"""Returns a mask function that is the intersection of provided mask functions"""
if not all(callable(arg) for arg in mask_functions):
raise RuntimeError(f"All inputs should be callable mask_functions: {mask_functions}")
def and_mask(batch_idx, head_idx, q_idx, kv_idx):
result = q_idx.new_ones((), dtype=torch.bool)
for mask in mask_functions:
result = result & mask(batch_idx, head_idx, q_idx, kv_idx).to(result.device)
return result
return and_mask
def causal_mask_function(batch_idx: int, head_idx: int, q_idx: int, kv_idx: int) -> bool:
"""
This creates a basic lower-diagonal causal mask.
"""
return kv_idx <= q_idx
def padding_mask_function(padding_mask: torch.Tensor) -> Callable:
"""
This return the mask_function function corresponding to a 2D padding mask.
"""
def inner_mask(batch_idx: int, head_idx: int, q_idx: int, kv_idx: int) -> bool:
# Note that here the mask should ALWAYS be at least of the max `kv_index` size in the dimension 1. This is because
# we cannot pad it here in the mask_function as we don't know the final size, and we cannot try/except, as it is not
# vectorizable on accelerator devices
return padding_mask[batch_idx, kv_idx]
return inner_mask
def prepare_padding_mask(
attention_mask: Optional[torch.Tensor], kv_length: int, kv_offset: int
) -> Optional[torch.Tensor]:
"""
From the 2D attention mask, prepare the correct padding mask to use by potentially padding it.
"""
local_padding_mask = attention_mask
if attention_mask is not None:
# Pad it if necessary
if (padding_length := kv_length + kv_offset - attention_mask.shape[-1]) > 0:
local_padding_mask = torch.nn.functional.pad(attention_mask, (0, padding_length))
return local_padding_mask
def _non_vmap_expansion_sdpa(
batch_indices: torch.Tensor, head_indices: torch.Tensor, q_indices: torch.Tensor, kv_indices: torch.Tensor
):
"""
Used to broadcast our mask_functions over the all 4 dimensions (b_idx, h_idx, q_idx, kv_idx) of the inputs.
Allows the usage of any index-based mask function without relying on vmap.
NOTE: This is limited to index based functions only and is not guaranteed to work otherwise.
Reference:
- https://github.com/huggingface/optimum-onnx/blob/c123e8f4fab61b54a8e0e31ce74462bcacca576e/optimum/exporters/onnx/model_patcher.py#L362-L365
"""
batch_indices = batch_indices[:, None, None, None]
head_indices = head_indices[None, :, None, None]
q_indices = q_indices[None, None, :, None]
kv_indices = kv_indices[None, None, None, :]
return batch_indices, head_indices, q_indices, kv_indices
def sdpa_mask(
batch_size: int,
cache_position: torch.Tensor,
kv_length: int,
kv_offset: int = 0,
mask_function: Callable = causal_mask_function,
attention_mask: Optional[torch.Tensor] = None,
local_size: Optional[int] = None,
allow_is_causal_skip: bool = True,
allow_is_bidirectional_skip: bool = False,
allow_torch_fix: bool = True,
use_vmap: bool = False,
**kwargs,
) -> Optional[torch.Tensor]:
"""
Create a 4D boolean mask of shape `(batch_size, 1, query_length, kv_length)` where a value of True indicates that
the element should take part in the attention computation, and False that it should not.
This function can only be used with torch>=2.5, as the context manager is otherwise not available.
Args:
batch_size (`int`):
The batch size of the input sequence.
cache_position (`torch.Tensor`):
A tensor of shape (query_length,) indicating the current indices of the input sequence elements.
kv_length (`int`):
The size that the key and value states will have during the attention computation.
kv_offset (`int`, optional):
An optional offset to indicate at which first position the key and values states will refer to.
mask_function (`Callable`):
The mask factory function describing the mask pattern.
attention_mask (`torch.Tensor`, optional):
The 2D attention mask corresponding to padded tokens of shape (batch_size, number_of_seen_tokens+q_length)
local_size (`int`, optional):
The size of the local attention, if we do not use full attention. This is used only if `allow_is_causal_skip=True`
to try to skip mask creation if possible.
allow_is_causal_skip (`bool`, optional):
Whether to allow to return `None` for the mask under conditions where we can use the `is_causal` argument in
`torch.sdpa` instead. Default to `True`.
allow_is_bidirectional_skip (`bool`, optional):
Whether to allow to return `None` for the mask under conditions where we do not have to add any bias,
i.e. full attention without any padding. Default to `False`.
allow_torch_fix (`bool`, optional):
Whether to update the mask in case a query is not attending to any tokens, to solve a bug in torch's older
versions. We need an arg to skip it when using eager. By default `True`.
use_vmap (`bool`, optional):
Whether to use `vmap` during the mask construction or not. Allows powerful custom patterns that may not be
index-based (for the cost of speed performance). By default `False`.
## Creating a simple causal mask:
To create the following causal mask:
0 ■ ⬚ ⬚ ⬚ ⬚
1 ■ ■ ⬚ ⬚ ⬚
2 ■ ■ ■ ⬚ ⬚
3 ■ ■ ■ ■ ⬚
4 ■ ■ ■ ■ ■
You can do
```python
>>> sdpa_mask(batch_size=1, cache_position=torch.arange(5), kv_length=5)
>>> tensor([[[[ True, False, False, False, False],
[ True, True, False, False, False],
[ True, True, True, False, False],
[ True, True, True, True, False],
[ True, True, True, True, True]]]])
```
## Creating a sliding window mask:
To create the following sliding window mask (`sliding_window=3`):
0 ■ ⬚ ⬚ ⬚ ⬚
1 ■ ■ ⬚ ⬚ ⬚
2 ■ ■ ■ ⬚ ⬚
3 ⬚ ■ ■ ■ ⬚
4 ⬚ ⬚ ■ ■ ■
You can do
```python
>>> sdpa_mask(batch_size=1, cache_position=torch.arange(5), kv_length=5, mask_function=sliding_window_causal_mask_function(3))
>>> tensor([[[[ True, False, False, False, False],
[ True, True, False, False, False],
[ True, True, True, False, False],
[False, True, True, True, False],
[False, False, True, True, True]]]])
```
## Creating a chunked attention mask
To create the following chunked attention mask (`chunk_size=3`):
0 ■ ⬚ ⬚ ⬚ ⬚
1 ■ ■ ⬚ ⬚ ⬚
2 ■ ■ ■ ⬚ ⬚
3 ⬚ ⬚ ⬚ ■ ⬚
4 ⬚ ⬚ ⬚ ■ ■
You can do
```python
>>> sdpa_mask(batch_size=1, cache_position=torch.arange(5), kv_length=5, mask_function=chunked_causal_mask_function(3, torch.zeros(1, dtype=int)))
>>> tensor([[[[ True, False, False, False, False],
[ True, True, False, False, False],
[ True, True, True, False, False],
[False, False, False, True, False],
[False, False, False, True, True]]]])
```
"""
q_length = cache_position.shape[0]
# Potentially pad the 2D mask
padding_mask = prepare_padding_mask(attention_mask, kv_length, kv_offset)
# Potentially add the padding 2D mask
if padding_mask is not None:
mask_function = and_masks(mask_function, padding_mask_function(padding_mask))
batch_arange = torch.arange(batch_size, device=cache_position.device)
head_arange = torch.arange(1, device=cache_position.device)
# Similar to `kv_arange = torch.arange(start=kv_offset, end=kv_offset + kv_length, device=cache_position.device)`
# but without data-dependent slicing (i.e. torch.compile friendly)
kv_arange = torch.arange(kv_length, device=cache_position.device) + kv_offset
# Actual mask creation
# Apply mask function element-wise through broadcasting
attention_mask = mask_function(*_non_vmap_expansion_sdpa(batch_arange, head_arange, cache_position, kv_arange))
# Expand the mask to match batch size and query length if they weren't used in the mask function
attention_mask = attention_mask.expand(batch_size, -1, q_length, kv_length)
return attention_mask
+7
View File
@@ -24,6 +24,8 @@ logger = init_logger(__name__)
_TEXT_TO_VIDEO_DIT_MODELS = {
"HunyuanVideoTransformer3DModel":
("dits", "hunyuanvideo", "HunyuanVideoTransformer3DModel"),
"HunyuanVideo15Transformer3DModel":
("dits", "hunyuanvideo15", "HunyuanVideo15Transformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"StepVideoModel": ("dits", "stepvideo", "StepVideoModel"),
@@ -34,6 +36,8 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
# "HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoDiT"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"MatrixGameWanModel": ("dits", "matrix_game", "MatrixGameWanModel"),
"CausalMatrixGameWanModel": ("dits", "matrix_game", "CausalMatrixGameWanModel"),
}
_TEXT_ENCODER_MODELS = {
@@ -43,16 +47,19 @@ _TEXT_ENCODER_MODELS = {
"T5EncoderModel": ("encoders", "t5", "T5EncoderModel"),
"STEP1TextEncoder": ("encoders", "stepllm", "STEP1TextEncoder"),
"BertModel": ("encoders", "clip", "CLIPTextModel"),
"Qwen2_5_VLTextModel": ("encoders", "qwen2_5", "Qwen2_5_VLTextModel"),
}
_IMAGE_ENCODER_MODELS: dict[str, tuple] = {
# "HunyuanVideoTransformer3DModel": ("image_encoder", "hunyuanvideo", "HunyuanVideoImageEncoder"),
"CLIPVisionModelWithProjection": ("encoders", "clip", "CLIPVisionModel"),
"CLIPVisionModel": ("encoders", "clip", "CLIPVisionModel"),
}
_VAE_MODELS = {
"AutoencoderKLHunyuanVideo":
("vaes", "hunyuanvae", "AutoencoderKLHunyuanVideo"),
"AutoencoderKLHunyuanVideo15": ("vaes", "hunyuan15vae", "AutoencoderKLHunyuanVideo15"),
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo")
}
+703
View File
@@ -0,0 +1,703 @@
# SPDX-License-Identifier: Apache-2.0
# Adapted from diffusers
# Copyright 2025 The Hunyuan Team and The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
from typing import Optional, Tuple, Union
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
import torch.utils.checkpoint
from fastvideo.layers.activation import get_act_fn
from fastvideo.configs.models.vaes import Hunyuan15VAEConfig
from fastvideo.models.vaes.common import ParallelTiledVAE
class HunyuanVideo15CausalConv3d(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
kernel_size: Union[int, Tuple[int, int, int]] = 3,
stride: Union[int, Tuple[int, int, int]] = 1,
padding: Union[int, Tuple[int, int, int]] = 0,
dilation: Union[int, Tuple[int, int, int]] = 1,
bias: bool = True,
pad_mode: str = "replicate",
) -> None:
super().__init__()
kernel_size = (kernel_size, kernel_size, kernel_size) if isinstance(kernel_size, int) else kernel_size
self.pad_mode = pad_mode
self.time_causal_padding = (
kernel_size[0] // 2,
kernel_size[0] // 2,
kernel_size[1] // 2,
kernel_size[1] // 2,
kernel_size[2] - 1,
0,
)
self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride, padding, dilation, bias=bias)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = F.pad(hidden_states, self.time_causal_padding, mode=self.pad_mode)
return self.conv(hidden_states)
class HunyuanVideo15RMS_norm(nn.Module):
r"""
A custom RMS normalization layer.
Args:
dim (int): The number of dimensions to normalize over.
channel_first (bool, optional): Whether the input tensor has channels as the first dimension.
Default is True.
images (bool, optional): Whether the input represents image data. Default is True.
bias (bool, optional): Whether to include a learnable bias term. Default is False.
"""
def __init__(self, dim: int, channel_first: bool = True, images: bool = True, bias: bool = False) -> None:
super().__init__()
broadcastable_dims = (1, 1, 1) if not images else (1, 1)
shape = (dim, *broadcastable_dims) if channel_first else (dim,)
self.channel_first = channel_first
self.scale = dim**0.5
self.gamma = nn.Parameter(torch.ones(shape))
self.bias = nn.Parameter(torch.zeros(shape)) if bias else 0.0
def forward(self, x):
return F.normalize(x, dim=(1 if self.channel_first else -1)) * self.scale * self.gamma + self.bias
class HunyuanVideo15AttnBlock(nn.Module):
def __init__(self, in_channels: int):
super().__init__()
self.in_channels = in_channels
self.norm = HunyuanVideo15RMS_norm(in_channels, images=False)
self.to_q = nn.Conv3d(in_channels, in_channels, kernel_size=1)
self.to_k = nn.Conv3d(in_channels, in_channels, kernel_size=1)
self.to_v = nn.Conv3d(in_channels, in_channels, kernel_size=1)
self.proj_out = nn.Conv3d(in_channels, in_channels, kernel_size=1)
@staticmethod
def prepare_causal_attention_mask(n_frame: int, n_hw: int, dtype, device, batch_size: int = None):
"""Prepare a causal attention mask for 3D videos.
Args:
n_frame (int): Number of frames (temporal length).
n_hw (int): Product of height and width.
dtype: Desired mask dtype.
device: Device for the mask.
batch_size (int, optional): If set, expands for batch.
Returns:
torch.Tensor: Causal attention mask.
"""
seq_len = n_frame * n_hw
mask = torch.full((seq_len, seq_len), float("-inf"), dtype=dtype, device=device)
for i in range(seq_len):
i_frame = i // n_hw
mask[i, : (i_frame + 1) * n_hw] = 0
if batch_size is not None:
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
return mask
def forward(self, x: torch.Tensor) -> torch.Tensor:
identity = x
x = self.norm(x)
query = self.to_q(x)
key = self.to_k(x)
value = self.to_v(x)
batch_size, channels, frames, height, width = query.shape
query = query.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous()
key = key.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous()
value = value.reshape(batch_size, channels, frames * height * width).permute(0, 2, 1).unsqueeze(1).contiguous()
attention_mask = self.prepare_causal_attention_mask(
frames, height * width, query.dtype, query.device, batch_size=batch_size
)
x = nn.functional.scaled_dot_product_attention(query, key, value, attn_mask=attention_mask)
# batch_size, 1, frames * height * width, channels
x = x.squeeze(1).reshape(batch_size, frames, height, width, channels).permute(0, 4, 1, 2, 3)
x = self.proj_out(x)
return x + identity
class HunyuanVideo15Upsample(nn.Module):
def __init__(self, in_channels: int, out_channels: int, add_temporal_upsample: bool = True):
super().__init__()
factor = 2 * 2 * 2 if add_temporal_upsample else 1 * 2 * 2
self.conv = HunyuanVideo15CausalConv3d(in_channels, out_channels * factor, kernel_size=3)
self.add_temporal_upsample = add_temporal_upsample
self.repeats = factor * out_channels // in_channels
@staticmethod
def _dcae_upsample_rearrange(tensor, r1=1, r2=2, r3=2):
"""
Convert (b, r1*r2*r3*c, f, h, w) -> (b, c, r1*f, r2*h, r3*w)
Args:
tensor: Input tensor of shape (b, r1*r2*r3*c, f, h, w)
r1: temporal upsampling factor
r2: height upsampling factor
r3: width upsampling factor
"""
b, packed_c, f, h, w = tensor.shape
factor = r1 * r2 * r3
c = packed_c // factor
tensor = tensor.view(b, r1, r2, r3, c, f, h, w)
tensor = tensor.permute(0, 4, 5, 1, 6, 2, 7, 3)
return tensor.reshape(b, c, f * r1, h * r2, w * r3)
def forward(self, x: torch.Tensor):
r1 = 2 if self.add_temporal_upsample else 1
h = self.conv(x)
if self.add_temporal_upsample:
h_first = h[:, :, :1, :, :]
h_first = self._dcae_upsample_rearrange(h_first, r1=1, r2=2, r3=2)
h_first = h_first[:, : h_first.shape[1] // 2]
h_next = h[:, :, 1:, :, :]
h_next = self._dcae_upsample_rearrange(h_next, r1=r1, r2=2, r3=2)
h = torch.cat([h_first, h_next], dim=2)
# shortcut computation
x_first = x[:, :, :1, :, :]
x_first = self._dcae_upsample_rearrange(x_first, r1=1, r2=2, r3=2)
x_first = x_first.repeat_interleave(repeats=self.repeats // 2, dim=1)
x_next = x[:, :, 1:, :, :]
x_next = self._dcae_upsample_rearrange(x_next, r1=r1, r2=2, r3=2)
x_next = x_next.repeat_interleave(repeats=self.repeats, dim=1)
shortcut = torch.cat([x_first, x_next], dim=2)
else:
h = self._dcae_upsample_rearrange(h, r1=r1, r2=2, r3=2)
shortcut = x.repeat_interleave(repeats=self.repeats, dim=1)
shortcut = self._dcae_upsample_rearrange(shortcut, r1=r1, r2=2, r3=2)
return h + shortcut
class HunyuanVideo15Downsample(nn.Module):
def __init__(self, in_channels: int, out_channels: int, add_temporal_downsample: bool = True):
super().__init__()
factor = 2 * 2 * 2 if add_temporal_downsample else 1 * 2 * 2
self.conv = HunyuanVideo15CausalConv3d(in_channels, out_channels // factor, kernel_size=3)
self.add_temporal_downsample = add_temporal_downsample
self.group_size = factor * in_channels // out_channels
@staticmethod
def _dcae_downsample_rearrange(tensor, r1=1, r2=2, r3=2):
"""
Convert (b, c, r1*f, r2*h, r3*w) -> (b, r1*r2*r3*c, f, h, w)
This packs spatial/temporal dimensions into channels (opposite of upsample)
"""
b, c, packed_f, packed_h, packed_w = tensor.shape
f, h, w = packed_f // r1, packed_h // r2, packed_w // r3
tensor = tensor.view(b, c, f, r1, h, r2, w, r3)
tensor = tensor.permute(0, 3, 5, 7, 1, 2, 4, 6)
return tensor.reshape(b, r1 * r2 * r3 * c, f, h, w)
def forward(self, x: torch.Tensor):
r1 = 2 if self.add_temporal_downsample else 1
h = self.conv(x)
if self.add_temporal_downsample:
h_first = h[:, :, :1, :, :]
h_first = self._dcae_downsample_rearrange(h_first, r1=1, r2=2, r3=2)
h_first = torch.cat([h_first, h_first], dim=1)
h_next = h[:, :, 1:, :, :]
h_next = self._dcae_downsample_rearrange(h_next, r1=r1, r2=2, r3=2)
h = torch.cat([h_first, h_next], dim=2)
# shortcut computation
x_first = x[:, :, :1, :, :]
x_first = self._dcae_downsample_rearrange(x_first, r1=1, r2=2, r3=2)
B, C, T, H, W = x_first.shape
x_first = x_first.view(B, h.shape[1], self.group_size // 2, T, H, W).mean(dim=2)
x_next = x[:, :, 1:, :, :]
x_next = self._dcae_downsample_rearrange(x_next, r1=r1, r2=2, r3=2)
B, C, T, H, W = x_next.shape
x_next = x_next.view(B, h.shape[1], self.group_size, T, H, W).mean(dim=2)
shortcut = torch.cat([x_first, x_next], dim=2)
else:
h = self._dcae_downsample_rearrange(h, r1=r1, r2=2, r3=2)
shortcut = self._dcae_downsample_rearrange(x, r1=r1, r2=2, r3=2)
B, C, T, H, W = shortcut.shape
shortcut = shortcut.view(B, h.shape[1], self.group_size, T, H, W).mean(dim=2)
return h + shortcut
class HunyuanVideo15ResnetBlock(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: Optional[int] = None,
non_linearity: str = "swish",
) -> None:
super().__init__()
out_channels = out_channels or in_channels
self.nonlinearity = get_act_fn(non_linearity)
self.norm1 = HunyuanVideo15RMS_norm(in_channels, images=False)
self.conv1 = HunyuanVideo15CausalConv3d(in_channels, out_channels, kernel_size=3)
self.norm2 = HunyuanVideo15RMS_norm(out_channels, images=False)
self.conv2 = HunyuanVideo15CausalConv3d(out_channels, out_channels, kernel_size=3)
self.conv_shortcut = None
if in_channels != out_channels:
self.conv_shortcut = nn.Conv3d(in_channels, out_channels, kernel_size=1, stride=1, padding=0)
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
residual = hidden_states
hidden_states = self.norm1(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.conv1(hidden_states)
hidden_states = self.norm2(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.conv2(hidden_states)
if self.conv_shortcut is not None:
residual = self.conv_shortcut(residual)
return hidden_states + residual
class HunyuanVideo15MidBlock(nn.Module):
def __init__(
self,
in_channels: int,
num_layers: int = 1,
add_attention: bool = True,
) -> None:
super().__init__()
self.add_attention = add_attention
# There is always at least one resnet
resnets = [
HunyuanVideo15ResnetBlock(
in_channels=in_channels,
out_channels=in_channels,
)
]
attentions = []
for _ in range(num_layers):
if self.add_attention:
attentions.append(HunyuanVideo15AttnBlock(in_channels))
else:
attentions.append(None)
resnets.append(
HunyuanVideo15ResnetBlock(
in_channels=in_channels,
out_channels=in_channels,
)
)
self.attentions = nn.ModuleList(attentions)
self.resnets = nn.ModuleList(resnets)
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.resnets[0](hidden_states)
for attn, resnet in zip(self.attentions, self.resnets[1:]):
if attn is not None:
hidden_states = attn(hidden_states)
hidden_states = resnet(hidden_states)
return hidden_states
class HunyuanVideo15DownBlock3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
num_layers: int = 1,
downsample_out_channels: Optional[int] = None,
add_temporal_downsample: int = True,
) -> None:
super().__init__()
resnets = []
for i in range(num_layers):
in_channels = in_channels if i == 0 else out_channels
resnets.append(
HunyuanVideo15ResnetBlock(
in_channels=in_channels,
out_channels=out_channels,
)
)
self.resnets = nn.ModuleList(resnets)
if downsample_out_channels is not None:
self.downsamplers = nn.ModuleList(
[
HunyuanVideo15Downsample(
out_channels,
out_channels=downsample_out_channels,
add_temporal_downsample=add_temporal_downsample,
)
]
)
else:
self.downsamplers = None
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
for resnet in self.resnets:
hidden_states = resnet(hidden_states)
if self.downsamplers is not None:
for downsampler in self.downsamplers:
hidden_states = downsampler(hidden_states)
return hidden_states
class HunyuanVideo15UpBlock3D(nn.Module):
def __init__(
self,
in_channels: int,
out_channels: int,
num_layers: int = 1,
upsample_out_channels: Optional[int] = None,
add_temporal_upsample: bool = True,
) -> None:
super().__init__()
resnets = []
for i in range(num_layers):
input_channels = in_channels if i == 0 else out_channels
resnets.append(
HunyuanVideo15ResnetBlock(
in_channels=input_channels,
out_channels=out_channels,
)
)
self.resnets = nn.ModuleList(resnets)
if upsample_out_channels is not None:
self.upsamplers = nn.ModuleList(
[
HunyuanVideo15Upsample(
out_channels,
out_channels=upsample_out_channels,
add_temporal_upsample=add_temporal_upsample,
)
]
)
else:
self.upsamplers = None
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
if torch.is_grad_enabled() and self.gradient_checkpointing:
for resnet in self.resnets:
hidden_states = self._gradient_checkpointing_func(resnet, hidden_states)
else:
for resnet in self.resnets:
hidden_states = resnet(hidden_states)
if self.upsamplers is not None:
for upsampler in self.upsamplers:
hidden_states = upsampler(hidden_states)
return hidden_states
class HunyuanVideo15Encoder3D(nn.Module):
r"""
3D vae encoder for HunyuanImageRefiner.
"""
def __init__(
self,
in_channels: int = 3,
out_channels: int = 64,
block_out_channels: Tuple[int, ...] = (128, 256, 512, 1024, 1024),
layers_per_block: int = 2,
temporal_compression_ratio: int = 4,
spatial_compression_ratio: int = 16,
downsample_match_channel: bool = True,
) -> None:
super().__init__()
self.in_channels = in_channels
self.out_channels = out_channels
self.group_size = block_out_channels[-1] // self.out_channels
self.conv_in = HunyuanVideo15CausalConv3d(in_channels, block_out_channels[0], kernel_size=3)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
input_channel = block_out_channels[0]
for i in range(len(block_out_channels)):
add_spatial_downsample = i < np.log2(spatial_compression_ratio)
output_channel = block_out_channels[i]
if not add_spatial_downsample:
down_block = HunyuanVideo15DownBlock3D(
num_layers=layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
downsample_out_channels=None,
add_temporal_downsample=False,
)
input_channel = output_channel
else:
add_temporal_downsample = i >= np.log2(spatial_compression_ratio // temporal_compression_ratio)
downsample_out_channels = block_out_channels[i + 1] if downsample_match_channel else output_channel
down_block = HunyuanVideo15DownBlock3D(
num_layers=layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
downsample_out_channels=downsample_out_channels,
add_temporal_downsample=add_temporal_downsample,
)
input_channel = downsample_out_channels
self.down_blocks.append(down_block)
self.mid_block = HunyuanVideo15MidBlock(in_channels=block_out_channels[-1])
self.norm_out = HunyuanVideo15RMS_norm(block_out_channels[-1], images=False)
self.conv_act = nn.SiLU()
self.conv_out = HunyuanVideo15CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.conv_in(hidden_states)
if torch.is_grad_enabled() and self.gradient_checkpointing:
for down_block in self.down_blocks:
hidden_states = self._gradient_checkpointing_func(down_block, hidden_states)
hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states)
else:
for down_block in self.down_blocks:
hidden_states = down_block(hidden_states)
hidden_states = self.mid_block(hidden_states)
batch_size, _, frame, height, width = hidden_states.shape
short_cut = hidden_states.view(batch_size, -1, self.group_size, frame, height, width).mean(dim=2)
hidden_states = self.norm_out(hidden_states)
hidden_states = self.conv_act(hidden_states)
hidden_states = self.conv_out(hidden_states)
hidden_states += short_cut
return hidden_states
class HunyuanVideo15Decoder3D(nn.Module):
r"""
Causal decoder for 3D video-like data used for HunyuanImage-1.5 Refiner.
"""
def __init__(
self,
in_channels: int = 32,
out_channels: int = 3,
block_out_channels: Tuple[int, ...] = (1024, 1024, 512, 256, 128),
layers_per_block: int = 2,
spatial_compression_ratio: int = 16,
temporal_compression_ratio: int = 4,
upsample_match_channel: bool = True,
):
super().__init__()
self.layers_per_block = layers_per_block
self.in_channels = in_channels
self.out_channels = out_channels
self.repeat = block_out_channels[0] // self.in_channels
self.conv_in = HunyuanVideo15CausalConv3d(self.in_channels, block_out_channels[0], kernel_size=3)
self.up_blocks = nn.ModuleList([])
# mid
self.mid_block = HunyuanVideo15MidBlock(in_channels=block_out_channels[0])
# up
input_channel = block_out_channels[0]
for i in range(len(block_out_channels)):
output_channel = block_out_channels[i]
add_spatial_upsample = i < np.log2(spatial_compression_ratio)
add_temporal_upsample = i < np.log2(temporal_compression_ratio)
if add_spatial_upsample or add_temporal_upsample:
upsample_out_channels = block_out_channels[i + 1] if upsample_match_channel else output_channel
up_block = HunyuanVideo15UpBlock3D(
num_layers=self.layers_per_block + 1,
in_channels=input_channel,
out_channels=output_channel,
upsample_out_channels=upsample_out_channels,
add_temporal_upsample=add_temporal_upsample,
)
input_channel = upsample_out_channels
else:
up_block = HunyuanVideo15UpBlock3D(
num_layers=self.layers_per_block + 1,
in_channels=input_channel,
out_channels=output_channel,
upsample_out_channels=None,
add_temporal_upsample=False,
)
input_channel = output_channel
self.up_blocks.append(up_block)
# out
self.norm_out = HunyuanVideo15RMS_norm(block_out_channels[-1], images=False)
self.conv_act = nn.SiLU()
self.conv_out = HunyuanVideo15CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
self.gradient_checkpointing = False
def forward(self, hidden_states: torch.Tensor) -> torch.Tensor:
hidden_states = self.conv_in(hidden_states) + hidden_states.repeat_interleave(repeats=self.repeat, dim=1)
if torch.is_grad_enabled() and self.gradient_checkpointing:
hidden_states = self._gradient_checkpointing_func(self.mid_block, hidden_states)
for up_block in self.up_blocks:
hidden_states = self._gradient_checkpointing_func(up_block, hidden_states)
else:
hidden_states = self.mid_block(hidden_states)
for up_block in self.up_blocks:
hidden_states = up_block(hidden_states)
# post-process
hidden_states = self.norm_out(hidden_states)
hidden_states = self.conv_act(hidden_states)
hidden_states = self.conv_out(hidden_states)
return hidden_states
class AutoencoderKLHunyuanVideo15(nn.Module, ParallelTiledVAE):
r"""
A VAE model with KL loss for encoding videos into latents and decoding latent representations into videos. Used for
HunyuanVideo-1.5.
This model inherits from [`ModelMixin`]. Check the superclass documentation for it's generic methods implemented
for all models (such as downloading or saving).
"""
_supports_gradient_checkpointing = True
def __init__(
self,
config: Hunyuan15VAEConfig,
) -> None:
nn.Module.__init__(self)
ParallelTiledVAE.__init__(self, config)
if config.load_encoder:
self.encoder = HunyuanVideo15Encoder3D(
in_channels=config.in_channels,
out_channels=config.latent_channels * 2,
block_out_channels=config.block_out_channels,
layers_per_block=config.layers_per_block,
temporal_compression_ratio=config.temporal_compression_ratio,
spatial_compression_ratio=config.spatial_compression_ratio,
downsample_match_channel=config.downsample_match_channel,
)
if config.load_decoder:
self.decoder = HunyuanVideo15Decoder3D(
in_channels=config.latent_channels,
out_channels=config.out_channels,
block_out_channels=list(reversed(config.block_out_channels)),
layers_per_block=config.layers_per_block,
temporal_compression_ratio=config.temporal_compression_ratio,
spatial_compression_ratio=config.spatial_compression_ratio,
upsample_match_channel=config.upsample_match_channel,
)
# When decoding spatially large video latents, the memory requirement is very high. By breaking the video latent
# frames spatially into smaller tiles and performing multiple forward passes for decoding, and then blending the
# intermediate tiles together, the memory requirement can be lowered.
self.use_tiling = False
# The minimal tile height and width for spatial tiling to be used
self.tile_sample_min_height = 256
self.tile_sample_min_width = 256
self.tile_sample_min_num_frames = 2000 # Fill in a random large number, as hy1.5 vae does not use temporal tiling
def _encode(self, x: torch.Tensor) -> torch.Tensor:
x = self.encoder(x)
return x
def _decode(self, z: torch.Tensor) -> torch.Tensor:
dec = self.decoder(z)
return dec
def forward(
self,
sample: torch.Tensor,
sample_posterior: bool = False,
return_dict: bool = True,
generator: Optional[torch.Generator] = None,
) -> torch.Tensor:
r"""
Args:
sample (`torch.Tensor`): Input sample.
sample_posterior (`bool`, *optional*, defaults to `False`):
Whether to sample from the posterior.
return_dict (`bool`, *optional*, defaults to `True`):
Whether or not to return a [`DecoderOutput`] instead of a plain tuple.
"""
x = sample
posterior = self.encode(x).latent_dist
if sample_posterior:
z = posterior.sample(generator=generator)
else:
z = posterior.mode()
dec = self.decode(z)
return dec
+6
View File
@@ -44,6 +44,12 @@ def build_pipeline(
config = verify_model_config_and_directory(model_path)
pipeline_name = config.get("_class_name")
if fastvideo_args.override_pipeline_cls_name:
logger.info("Overriding pipeline class name from %s to %s",
pipeline_name, fastvideo_args.override_pipeline_cls_name)
pipeline_name = fastvideo_args.override_pipeline_cls_name
if pipeline_name is None:
raise ValueError(
"Model config does not contain a _class_name attribute. "
@@ -0,0 +1,74 @@
# SPDX-License-Identifier: Apache-2.0
"""
Hunyuan video diffusion pipeline implementation.
This module contains an implementation of the Hunyuan video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage,
Hy15ImageEncodingStage)
# TODO(will): move PRECISION_TO_TYPE to better place
logger = init_logger(__name__)
class HunyuanVideo15Pipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae",
"transformer", "scheduler"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage_primary",
stage=TextEncodingStage(
text_encoders=[
self.get_module("text_encoder"),
self.get_module("text_encoder_2")
],
tokenizers=[
self.get_module("tokenizer"),
self.get_module("tokenizer_2")
],
))
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"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_encoding_stage",
stage=Hy15ImageEncodingStage(image_encoder=None,
image_processor=None))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = HunyuanVideo15Pipeline
@@ -0,0 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.pipelines.basic.matrixgame.matrixgame_i2v_pipeline import (
MatrixGamePipeline)
from fastvideo.pipelines.basic.matrixgame.matrixgame_causal_dmd_pipeline import (
MatrixGameCausalDMDPipeline)
__all__ = ["MatrixGamePipeline", "MatrixGameCausalDMDPipeline"]
@@ -0,0 +1,73 @@
# SPDX-License-Identifier: Apache-2.0
"""Matrix-Game causal DMD pipeline implementation."""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.stages import (ConditioningStage, DecodingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
MatrixGameImageEncodingStage,
MatrixGameCausalDenoisingStage)
from fastvideo.pipelines.stages.image_encoding import (
MatrixGameImageVAEEncodingStage)
logger = init_logger(__name__)
class MatrixGameCausalDMDPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"vae", "transformer", "scheduler", "image_encoder", "image_processor"
]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
if (self.get_module("text_encoder", None) is not None
and self.get_module("tokenizer", None) is not None):
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
if (self.get_module("image_encoder", None) is not None
and self.get_module("image_processor", None) is not None):
self.add_stage(
stage_name="image_encoding_stage",
stage=MatrixGameImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
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="image_latent_preparation_stage",
stage=MatrixGameImageVAEEncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=MatrixGameCausalDenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
pipeline=self,
vae=self.get_module("vae")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
logger.info(
"MatrixGameCausalDMDPipeline initialized with action support")
EntryClass = [MatrixGameCausalDMDPipeline]
@@ -0,0 +1,79 @@
# SPDX-License-Identifier: Apache-2.0
"""Matrix-Game I2V pipeline implementation."""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
from fastvideo.pipelines.stages import (MatrixGameImageEncodingStage,
ConditioningStage, DecodingStage,
DenoisingStage, InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
from fastvideo.pipelines.stages.image_encoding import (
MatrixGameImageVAEEncodingStage)
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
logger = init_logger(__name__)
class MatrixGamePipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"vae", "transformer", "scheduler", "image_encoder", "image_processor"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
if (self.get_module("text_encoder", None) is not None
and self.get_module("tokenizer", None) is not None):
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
if (self.get_module("image_encoder", None) is not None
and self.get_module("image_processor", None) is not None):
self.add_stage(
stage_name="image_encoding_stage",
stage=MatrixGameImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
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"),
transformer=self.get_module("transformer")))
self.add_stage(
stage_name="image_latent_preparation_stage",
stage=MatrixGameImageVAEEncodingStage(vae=self.get_module("vae")))
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")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = [MatrixGamePipeline]
+25 -7
View File
@@ -43,6 +43,9 @@ class ComposedPipelineBase(ABC):
training_args: TrainingArgs | None = None
fastvideo_args: FastVideoArgs | TrainingArgs | None = None
modules: dict[str, Any] = {}
# do not need to include moe related transformers
trainable_transformer_names: list[str] = ["transformer"]
trainable_transformer_modules: dict[str, torch.nn.Module] = {}
post_init_called: bool = False
# TODO(will): args should support both inference args and training args
@@ -87,13 +90,14 @@ class ComposedPipelineBase(ABC):
def set_trainable(self) -> None:
# Only train DiT
if getattr(self.fastvideo_args, "training_mode", False):
for name, module in self.modules.items():
for name, module in self.trainable_transformer_modules.items():
logger.info("Setting %s to requires_grad=True", name)
if not isinstance(module, torch.nn.Module):
logger.info(
"Skipping %s because it is not a torch.nn.Module", name)
continue
if "transformer" in name:
module.requires_grad_(True)
else:
module.requires_grad_(False)
module.requires_grad_(True)
module.train()
def post_init(self) -> None:
assert self.fastvideo_args is not None, "fastvideo_args must be set"
@@ -122,18 +126,32 @@ class ComposedPipelineBase(ABC):
fsdp_module_cls = FSDPModule
except Exception: # pragma: no cover - FSDP not always available
fsdp_module_cls = None
compile_kwargs = self.fastvideo_args.torch_compile_kwargs or {}
if fsdp_module_cls is not None and isinstance(
transformer_module, fsdp_module_cls):
logger.info(
"Transformer is already FSDP-wrapped; skipping torch.compile in pipeline"
)
else:
compile_kwargs = self.fastvideo_args.torch_compile_kwargs or {}
logger.info("Enabling torch.compile for DiT with kwargs=%s",
compile_kwargs)
self.modules["transformer"] = torch.compile(
transformer_module, **compile_kwargs)
logger.info("Torch Compile enabled for DiT")
if "transformer_2" in self.modules:
transformer_module_2 = self.modules["transformer_2"]
if fsdp_module_cls is not None and isinstance(
transformer_module_2, fsdp_module_cls):
logger.info(
"Transformer_2 is already FSDP-wrapped; skipping torch.compile in pipeline"
)
else:
logger.info(
"Enabling torch.compile for Transformer_2 with kwargs=%s",
compile_kwargs)
self.modules["transformer_2"] = torch.compile(
transformer_module_2, **compile_kwargs)
logger.info("Torch Compile enabled for DiT")
if not self.fastvideo_args.training_mode:
logger.info("Creating pipeline stages...")
+101 -72
View File
@@ -32,10 +32,9 @@ class LoRAPipeline(ComposedPipelineBase):
) # state dicts of loaded lora adapters (includes lora_A, lora_B, and lora_alpha)
cur_adapter_name: str = ""
cur_adapter_path: str = ""
lora_layers: dict[str, BaseLayerWithLoRA] = {}
lora_layers_critic: dict[str, BaseLayerWithLoRA] = {}
lora_layers: dict[str, dict[str, BaseLayerWithLoRA]] = {}
fastvideo_args: FastVideoArgs | TrainingArgs
exclude_lora_layers: list[str] = []
exclude_lora_layers: dict[str, list[str]] = {}
device: torch.device = get_local_torch_device()
lora_target_modules: list[str] | None = None
lora_path: str | None = None
@@ -47,8 +46,32 @@ class LoRAPipeline(ComposedPipelineBase):
def __init__(self, *args, **kwargs) -> None:
super().__init__(*args, **kwargs)
self.device = get_local_torch_device()
self.exclude_lora_layers = self.modules[
"transformer"].config.arch_config.exclude_lora_layers
# build list of trainable transformers
for transformer_name in self.trainable_transformer_names:
if transformer_name in self.modules and self.modules[
transformer_name] is not None:
self.trainable_transformer_modules[
transformer_name] = self.modules[transformer_name]
# check for transformer_2 in case of Wan2.2 MoE or fake_score_transformer_2
if transformer_name.endswith("_2"):
raise ValueError(
f"trainable_transformer_name override in pipelines should not include _2 suffix: {transformer_name}"
)
secondary_transformer_name = transformer_name + "_2"
if secondary_transformer_name in self.modules and self.modules[
secondary_transformer_name] is not None:
self.trainable_transformer_modules[
secondary_transformer_name] = self.modules[
secondary_transformer_name]
logger.info("trainable_transformer_modules: %s",
self.trainable_transformer_modules.keys())
for transformer_name, transformer_module in self.trainable_transformer_modules.items(
):
self.exclude_lora_layers[
transformer_name] = transformer_module.config.arch_config.exclude_lora_layers
self.lora_target_modules = self.fastvideo_args.lora_target_modules
self.lora_path = self.fastvideo_args.lora_path
self.lora_nickname = self.fastvideo_args.lora_nickname
@@ -65,8 +88,16 @@ class LoRAPipeline(ComposedPipelineBase):
if self.lora_target_modules is None:
self.lora_target_modules = [
"q_proj", "k_proj", "v_proj", "o_proj", "to_q", "to_k",
"to_v", "to_out", "to_qkv"
"to_v", "to_out", "to_qkv", "to_gate_compress"
]
logger.info(
"Using default lora_target_modules for all transformers: %s",
self.lora_target_modules)
else:
logger.warning(
"Using custom lora_target_modules for all transformers, which may not be intended: %s",
self.lora_target_modules)
self.convert_to_lora_layers()
# Inference
elif not self.training_mode and self.lora_path is not None:
@@ -100,13 +131,18 @@ class LoRAPipeline(ComposedPipelineBase):
super().set_trainable()
return
self.modules["transformer"].requires_grad_(False)
if "fake_score_transformer" in self.modules:
self.modules["fake_score_transformer"].requires_grad_(False)
device_mesh = init_device_mesh("cuda", (dist.get_world_size(), 1),
mesh_dim_names=["fake", "replicate"])
set_lora_grads(self.lora_layers, device_mesh)
set_lora_grads(self.lora_layers_critic, device_mesh)
for transformer_name, transformer_module in self.trainable_transformer_modules.items(
):
transformer_module.train()
transformer_module.requires_grad_(False)
if transformer_name in self.lora_layers:
set_lora_grads(self.lora_layers[transformer_name], device_mesh)
else:
raise ValueError(
f"Transformer {transformer_name} should be trainable but not found in lora_layers"
)
def convert_to_lora_layers(self) -> None:
"""
@@ -115,46 +151,33 @@ class LoRAPipeline(ComposedPipelineBase):
if self.lora_initialized:
return
self.lora_initialized = True
converted_count = 0
for name, layer in self.modules["transformer"].named_modules():
if not self.is_target_layer(name):
continue
excluded = False
for exclude_layer in self.exclude_lora_layers:
if exclude_layer in name:
excluded = True
break
if excluded:
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[name] = layer
replace_submodule(self.modules["transformer"], name, layer)
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():
for transformer_name, transformer_module in self.trainable_transformer_modules.items(
):
converted_count = 0
if transformer_name not in self.lora_layers:
self.lora_layers[transformer_name] = {}
logger.info("Converting %s to LoRA Transformer", transformer_name)
for name, layer in transformer_module.named_modules():
if not self.is_target_layer(name):
continue
excluded = False
for exclude_layer in self.exclude_lora_layers[transformer_name]:
if exclude_layer in name:
excluded = True
break
if excluded:
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)
self.lora_layers[transformer_name][name] = layer
replace_submodule(transformer_module, name, layer)
converted_count += 1
logger.info(
"Converted %d layers to LoRA layers in the critic model",
converted_count)
logger.info("Converted %d layers to LoRA layers", converted_count)
def set_lora_adapter(self,
lora_nickname: str,
@@ -238,39 +261,45 @@ class LoRAPipeline(ComposedPipelineBase):
# Merge the new adapter
adapted_count = 0
for name, layer in self.lora_layers.items():
lora_A_name = name + ".lora_A"
lora_B_name = name + ".lora_B"
lora_alpha_name = name + ".lora_alpha"
if lora_A_name in self.lora_adapters[lora_nickname]\
and lora_B_name in self.lora_adapters[lora_nickname]:
# Get alpha value for this layer (defaults to None if not present)
lora_A = self.lora_adapters[lora_nickname][lora_A_name]
lora_B = self.lora_adapters[lora_nickname][lora_B_name]
# Simple lookup - alpha stored with same naming scheme as lora_A/lora_B
alpha = self.lora_adapters[lora_nickname].get(
lora_alpha_name) if adapter_updated else None
for transformer_name, transformer_lora_layers in self.lora_layers.items(
):
for name, layer in transformer_lora_layers.items():
lora_A_name = name + ".lora_A"
lora_B_name = name + ".lora_B"
lora_alpha_name = name + ".lora_alpha"
if lora_A_name in self.lora_adapters[lora_nickname]\
and lora_B_name in self.lora_adapters[lora_nickname]:
# Get alpha value for this layer (defaults to None if not present)
lora_A = self.lora_adapters[lora_nickname][lora_A_name]
lora_B = self.lora_adapters[lora_nickname][lora_B_name]
# Simple lookup - alpha stored with same naming scheme as lora_A/lora_B
alpha = self.lora_adapters[lora_nickname].get(
lora_alpha_name) if adapter_updated else None
layer.set_lora_weights(
lora_A,
lora_B,
lora_alpha=alpha,
training_mode=self.fastvideo_args.training_mode,
lora_path=lora_path)
adapted_count += 1
else:
if rank == 0:
logger.warning(
"LoRA adapter %s does not contain the weights for layer %s. LoRA will not be applied to it.",
lora_path, name)
layer.disable_lora = True
layer.set_lora_weights(
lora_A,
lora_B,
lora_alpha=alpha,
training_mode=self.fastvideo_args.training_mode,
lora_path=lora_path)
adapted_count += 1
else:
if rank == 0:
logger.warning(
"LoRA adapter %s does not contain the weights for layer %s. LoRA will not be applied to it.",
lora_path, name)
layer.disable_lora = True
logger.info("Rank %d: LoRA adapter %s applied to %d layers", rank,
lora_path, adapted_count)
def merge_lora_weights(self) -> None:
for name, layer in self.lora_layers.items():
layer.merge_lora_weights()
for transformer_name, transformer_lora_layers in self.lora_layers.items(
):
for name, layer in transformer_lora_layers.items():
layer.merge_lora_weights()
def unmerge_lora_weights(self) -> None:
for name, layer in self.lora_layers.items():
layer.unmerge_lora_weights()
for transformer_name, transformer_lora_layers in self.lora_layers.items(
):
for name, layer in transformer_lora_layers.items():
layer.unmerge_lora_weights()
+7 -2
View File
@@ -115,10 +115,15 @@ class ForwardBatch:
# Latent tensors
latents: torch.Tensor | None = None
raw_latent_shape: torch.Tensor | None = None
raw_latent_shape: tuple[int, ...] | None = None
noise_pred: torch.Tensor | None = None
image_latent: torch.Tensor | None = None
# Action control inputs (Matrix-Game)
mouse_cond: torch.Tensor | None = None # Shape: (B, T, 2)
keyboard_cond: torch.Tensor | None = None # Shape: (B, T, K)
grid_sizes: torch.Tensor | None = None # Shape: (3,) [F,H,W]
# Latent dimensions
height_latents: list[int] | int | None = None
width_latents: list[int] | int | None = None
@@ -206,7 +211,7 @@ class TrainingBatch:
# Dataloader batch outputs
latents: torch.Tensor | None = None
raw_latent_shape: torch.Tensor | None = None
raw_latent_shape: tuple[int, ...] | None = None
noise_latents: torch.Tensor | None = None
encoder_hidden_states: torch.Tensor | None = None
encoder_attention_mask: torch.Tensor | None = None
+4 -1
View File
@@ -25,7 +25,10 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"WanCausalDMDPipeline": "wan",
"StepVideoPipeline": "stepvideo",
"HunyuanVideoPipeline": "hunyuan",
"Cosmos2VideoToWorldPipeline": "cosmos"
"HunyuanVideo15Pipeline": "hunyuan15",
"Cosmos2VideoToWorldPipeline": "cosmos",
"MatrixGamePipeline": "matrixgame",
"MatrixGameCausalDMDPipeline": "matrixgame",
}
_PREPROCESS_WORKLOAD_TYPE_TO_PIPELINE_NAME: dict[WorkloadType, str] = {
@@ -303,7 +303,7 @@ class BasePreprocessPipeline(ComposedPipelineBase):
if data is None:
continue
with torch.inference_mode():
with torch.no_grad():
# Filter out invalid samples (those with all zeros)
valid_indices = []
for i, pixel_values in enumerate(data["pixel_values"]):
@@ -37,7 +37,8 @@ from fastvideo.pipelines.stages import (DecodingStage, DenoisingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage)
TimestepPreparationStage,
Hy15ImageEncodingStage)
from fastvideo.utils import save_decoded_latents_as_video, shallow_asdict
logger = init_logger(__name__)
@@ -47,7 +48,7 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"""ODE Trajectory preprocessing pipeline implementation."""
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
"text_encoder", "text_encoder_2", "tokenizer", "tokenizer_2", "vae", "transformer", "scheduler"
]
preprocess_dataloader: StatefulDataLoader
preprocess_loader_iter: Iterator[dict[str, Any]]
@@ -61,19 +62,13 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
assert fastvideo_args.pipeline_config.flow_shift == 5
self.modules["scheduler"] = SelfForcingFlowMatchScheduler(
shift=fastvideo_args.pipeline_config.flow_shift,
sigma_min=0.0,
extra_one_step=True)
self.modules["scheduler"].set_timesteps(num_inference_steps=48,
denoising_strength=1.0)
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")],
text_encoders=[self.get_module("text_encoder"), self.get_module("text_encoder_2")],
tokenizers=[self.get_module("tokenizer"), self.get_module("tokenizer_2")],
))
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
@@ -82,6 +77,9 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None)))
self.add_stage(stage_name="image_encoding_stage",
stage=Hy15ImageEncodingStage(image_encoder=None,
image_processor=None))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
@@ -95,11 +93,13 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
args):
"""Preprocess text-only data and generate trajectory information."""
num_encoders = len(self.prompt_encoding_stage.text_encoders)
for batch_idx, data in enumerate(self.pbar):
if data is None:
continue
with torch.inference_mode():
with torch.no_grad():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
@@ -130,12 +130,9 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
prompt_embeds_list, prompt_masks_list = self.prompt_encoding_stage.encode_text(
batch_captions,
fastvideo_args,
encoder_index=[0],
encoder_index=list(range(num_encoders)),
return_attention_mask=True,
)
prompt_embeds = prompt_embeds_list[0]
prompt_attention_masks = prompt_masks_list[0]
assert prompt_embeds.shape[0] == prompt_attention_masks.shape[0]
sampling_params = SamplingParam.from_pretrained(args.model_path)
@@ -144,61 +141,52 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
negative_prompt_embeds_list, negative_prompt_masks_list = self.prompt_encoding_stage.encode_text(
sampling_params.negative_prompt,
fastvideo_args,
encoder_index=[0],
encoder_index=list(range(num_encoders)),
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
negative_prompt_embeds_list = []
negative_prompt_masks_list = []
trajectory_latents = []
trajectory_timesteps = []
trajectory_decoded = []
for i, (prompt_embed, prompt_attention_mask) in enumerate(
zip(prompt_embeds, prompt_attention_masks,
strict=False)):
prompt_embed = prompt_embed.unsqueeze(0)
prompt_attention_mask = prompt_attention_mask.unsqueeze(0)
# Collect the trajectory data (text-to-video generation)
batch = ForwardBatch(**shallow_asdict(sampling_params), )
batch.prompt_embeds = prompt_embeds_list
batch.prompt_attention_mask = prompt_masks_list
batch.negative_prompt_embeds = negative_prompt_embeds_list
batch.negative_attention_mask = negative_prompt_masks_list
batch.num_inference_steps = 50
batch.return_trajectory_latents = True
# Enabling this will save the decoded trajectory videos.
# Used for debugging.
batch.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
batch.num_frames = args.num_frames
batch.fps = args.train_fps
batch.guidance_scale = 6.0
batch.do_classifier_free_guidance = True
# Collect the trajectory data (text-to-video generation)
batch = ForwardBatch(**shallow_asdict(sampling_params), )
batch.prompt_embeds = [prompt_embed]
batch.prompt_attention_mask = [prompt_attention_mask]
batch.negative_prompt_embeds = [negative_prompt_embed]
batch.negative_attention_mask = [
negative_prompt_attention_mask
]
batch.num_inference_steps = 48
batch.return_trajectory_latents = True
# Enabling this will save the decoded trajectory videos.
# Used for debugging.
batch.return_trajectory_decoded = False
batch.height = args.max_height
batch.width = args.max_width
batch.fps = args.train_fps
batch.guidance_scale = 6.0
batch.do_classifier_free_guidance = True
result_batch = self.input_validation_stage(
batch, fastvideo_args)
result_batch = self.timestep_preparation_stage(
batch, fastvideo_args)
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.image_encoding_stage(result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch,
fastvideo_args)
result_batch = self.input_validation_stage(
batch, fastvideo_args)
result_batch = self.timestep_preparation_stage(
batch, fastvideo_args)
result_batch = self.latent_preparation_stage(
result_batch, fastvideo_args)
result_batch = self.denoising_stage(result_batch,
fastvideo_args)
result_batch = self.decoding_stage(result_batch,
fastvideo_args)
trajectory_latents.append(
result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(
result_batch.trajectory_timesteps.cpu())
trajectory_decoded.append(result_batch.trajectory_decoded)
trajectory_latents.append(
result_batch.trajectory_latents.cpu())
trajectory_timesteps.append(
result_batch.trajectory_timesteps.cpu())
trajectory_decoded.append(result_batch.trajectory_decoded)
# Prepare extra features for text-only processing
extra_features = {
@@ -209,10 +197,11 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
if batch.return_trajectory_decoded:
for i, decoded_frames in enumerate(trajectory_decoded):
for j, decoded_frame in enumerate(decoded_frames):
save_decoded_latents_as_video(
decoded_frame,
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
args.train_fps)
if j in [40, 44, 49]:
save_decoded_latents_as_video(
decoded_frame,
f"decoded_videos/trajectory_decoded_{i}_{j}.mp4",
args.train_fps)
# Prepare batch data for Parquet dataset
batch_data: list[dict[str, Any]] = []
@@ -227,7 +216,10 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
video_name = os.path.basename(video_path).split(".")[0]
# Convert tensors to numpy arrays
text_embedding = prompt_embeds[idx].cpu().numpy()
text_embedding = prompt_embeds_list[0].float().cpu().numpy()
text_mask = prompt_masks_list[0].cpu().numpy()
text_embedding_2 = prompt_embeds_list[1].float().cpu().numpy()
text_mask_2 = prompt_masks_list[1].cpu().numpy()
# Get extra features for this sample
sample_extra_features = {}
@@ -253,6 +245,9 @@ class PreprocessPipeline_ODE_Trajectory(BasePreprocessPipeline):
"trajectory_latents"],
trajectory_timesteps=sample_extra_features[
"trajectory_timesteps"],
text_embedding_2=text_embedding_2,
text_mask=text_mask,
text_mask_2=text_mask_2,
)
batch_data.append(record)
@@ -58,7 +58,7 @@ class PreprocessPipeline_Text(BasePreprocessPipeline):
if data is None:
continue
with torch.inference_mode():
with torch.no_grad():
# For text-only processing, we only need text data
# Filter out samples without text
valid_indices = []
+17 -16
View File
@@ -22,32 +22,33 @@ logger = init_logger(__name__)
def main(args) -> None:
args.model_path = maybe_download_model(args.model_path)
# args.model_path = maybe_download_model(args.model_path)
maybe_init_distributed_environment_and_model_parallel(1, 1)
num_gpus = int(os.environ["WORLD_SIZE"])
assert num_gpus == 1, "Only support 1 GPU"
pipeline_config = PipelineConfig.from_pretrained(args.model_path)
print(pipeline_config.__class__.__name__)
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)
# 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(),
dit_cpu_offload=False,
vae_cpu_offload=False,
text_encoder_cpu_offload=False,
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pipeline_config=pipeline_config,
)
if args.preprocess_task == "t2v":
@@ -5,6 +5,10 @@ from fastvideo.logger import init_logger
from fastvideo.utils import FlexibleArgumentParser
from fastvideo.workflow.workflow_base import WorkflowBase
import os
os.environ["MASTER_PORT"] = "29513"
logger = init_logger(__name__)
+8 -4
View File
@@ -14,13 +14,14 @@ from fastvideo.pipelines.stages.denoising import (CosmosDenoisingStage,
DenoisingStage,
DmdDenoisingStage)
from fastvideo.pipelines.stages.encoding import EncodingStage
from fastvideo.pipelines.stages.image_encoding import (ImageEncodingStage,
RefImageEncodingStage,
ImageVAEEncodingStage,
VideoVAEEncodingStage)
from fastvideo.pipelines.stages.image_encoding import (
ImageEncodingStage, MatrixGameImageEncodingStage, RefImageEncodingStage,
ImageVAEEncodingStage, VideoVAEEncodingStage, Hy15ImageEncodingStage)
from fastvideo.pipelines.stages.input_validation import InputValidationStage
from fastvideo.pipelines.stages.latent_preparation import (
CosmosLatentPreparationStage, LatentPreparationStage)
from fastvideo.pipelines.stages.matrixgame_denoising import (
MatrixGameCausalDenoisingStage)
from fastvideo.pipelines.stages.stepvideo_encoding import (
StepvideoPromptEncodingStage)
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
@@ -37,10 +38,13 @@ __all__ = [
"DenoisingStage",
"DmdDenoisingStage",
"CausalDMDDenosingStage",
"MatrixGameCausalDenoisingStage",
"CosmosDenoisingStage",
"EncodingStage",
"DecodingStage",
"ImageEncodingStage",
"MatrixGameImageEncodingStage",
"Hy15ImageEncodingStage",
"RefImageEncodingStage",
"ImageVAEEncodingStage",
"VideoVAEEncodingStage",
+18 -5
View File
@@ -77,19 +77,32 @@ class ConditioningStage(PipelineStage):
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify conditioning stage inputs."""
result = VerificationResult()
if not batch.prompt_embeds:
# No text encoder/prompt embeddings: skip checks and effectively disable CFG.
batch.do_classifier_free_guidance = False
return result
result.add_check("do_classifier_free_guidance",
batch.do_classifier_free_guidance, V.bool_value)
result.add_check("guidance_scale", batch.guidance_scale,
V.positive_float)
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
not batch.do_classifier_free_guidance or V.list_not_empty(x))
# Matrix-Game allow empty prompt
# embeddings when CFG isn't enabled.
if batch.do_classifier_free_guidance or batch.prompt_embeds:
result.add_check("prompt_embeds", batch.prompt_embeds,
V.list_not_empty)
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
def verify_output(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify conditioning stage outputs."""
result = VerificationResult()
result.add_check("prompt_embeds", batch.prompt_embeds, V.list_not_empty)
if batch.prompt_embeds is None or not batch.prompt_embeds:
batch.do_classifier_free_guidance = False
return result
if batch.do_classifier_free_guidance or batch.prompt_embeds:
result.add_check("prompt_embeds", batch.prompt_embeds,
V.list_not_empty)
return result
+27 -9
View File
@@ -76,21 +76,39 @@ class DecodingStage(PipelineStage):
vae_autocast_enabled = (
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
if hasattr(self.vae, 'scaling_factor'):
# denormalization for MatrixGame VAE
# z = z * std + mean during decode
if (hasattr(self.vae.config, 'latents_mean')
and hasattr(self.vae.config, 'latents_std')):
# Convert config values to tensors
latents_mean = torch.tensor(self.vae.config.latents_mean,
device=latents.device,
dtype=latents.dtype).view(
1, -1, 1, 1, 1)
latents_std = torch.tensor(self.vae.config.latents_std,
device=latents.device,
dtype=latents.dtype).view(
1, -1, 1, 1, 1)
# Apply denormalization: z = z * std + mean
latents = latents * latents_std + latents_mean
elif hasattr(self.vae, 'scaling_factor'):
# Standard VAE scaling
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
# 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",
+12 -57
View File
@@ -4,21 +4,16 @@ Denoising stage for diffusion pipelines.
"""
import inspect
import math
import weakref
from collections.abc import Iterable
from typing import Any
import torch
from einops import rearrange
from tqdm.auto import tqdm
from fastvideo.attention import get_attn_backend
from fastvideo.configs.pipelines.base import STA_Mode
from fastvideo.distributed import (get_local_torch_device, get_sp_parallel_rank,
get_sp_world_size, get_world_group)
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather)
from fastvideo.distributed import (get_local_torch_device, get_world_group)
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
@@ -130,22 +125,6 @@ class DenoisingStage(PipelineStage):
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
# Handle sequence parallelism if enabled
sp_world_size, rank_in_sp_group = get_sp_world_size(
), get_sp_parallel_rank()
sp_group = sp_world_size > 1
if sp_group:
latents = rearrange(batch.latents,
"b c (n t) h w -> b c n t h w",
n=sp_world_size).contiguous()
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
# Get timesteps and calculate warmup steps
timesteps = batch.timesteps
# TODO(will): remove this once we add input/output validation for stages
@@ -189,6 +168,14 @@ class DenoisingStage(PipelineStage):
},
)
action_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"mouse_cond": batch.mouse_cond,
"keyboard_cond": batch.keyboard_cond,
},
)
# Prepare STA parameters
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.prepare_sta_param(batch, fastvideo_args)
@@ -252,7 +239,6 @@ class DenoisingStage(PipelineStage):
1) * (batch.height // spatial_scale) * (
batch.width // spatial_scale) // (patch_size[1] *
patch_size[2])
seq_len = int(math.ceil(seq_len / sp_world_size)) * sp_world_size
# Initialize lists for ODE trajectory
trajectory_timesteps: list[torch.Tensor] = []
@@ -401,6 +387,7 @@ class DenoisingStage(PipelineStage):
guidance=guidance_expand,
**image_kwargs,
**pos_cond_kwargs,
**action_kwargs,
)
if batch.do_classifier_free_guidance:
@@ -417,6 +404,7 @@ class DenoisingStage(PipelineStage):
guidance=guidance_expand,
**image_kwargs,
**neg_cond_kwargs,
**action_kwargs,
)
noise_pred_text = noise_pred
@@ -464,15 +452,6 @@ class DenoisingStage(PipelineStage):
trajectory_tensor = None
trajectory_timesteps_tensor = None
# Gather results if using sequence parallelism
if sp_group:
latents = sequence_model_parallel_all_gather(latents, dim=2)
if batch.return_trajectory_latents:
trajectory_tensor = trajectory_tensor.to(
get_local_torch_device())
trajectory_tensor = sequence_model_parallel_all_gather(
trajectory_tensor, dim=3)
if trajectory_tensor is not None and trajectory_timesteps_tensor is not None:
batch.trajectory_timesteps = trajectory_timesteps_tensor.cpu()
batch.trajectory_latents = trajectory_tensor.cpu()
@@ -1095,23 +1074,6 @@ class DmdDenoisingStage(DenoisingStage):
dtype=torch.long,
device=get_local_torch_device())
# Handle sequence parallelism if enabled
sp_world_size, rank_in_sp_group = get_sp_world_size(
), get_sp_parallel_rank()
sp_group = sp_world_size > 1
if sp_group:
latents = rearrange(latents,
"b (n t) c h w -> b n t c h w",
n=sp_world_size).contiguous()
latents = latents[:, rank_in_sp_group, :, :, :, :]
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
# Run denoising loop
with self.progress_bar(total=len(timesteps)) as progress_bar:
for i, t in enumerate(timesteps):
@@ -1205,11 +1167,6 @@ class DmdDenoisingStage(DenoisingStage):
dtype=pred_video.dtype,
generator=batch.generator[0]).to(
self.device)
if sp_group:
noise = rearrange(noise,
"b (n t) c h w -> b n t c h w",
n=sp_world_size).contiguous()
noise = noise[:, rank_in_sp_group, :, :, :, :]
latents = self.scheduler.add_noise(
pred_video.flatten(0, 1), noise.flatten(0, 1),
next_timestep).unflatten(0, pred_video.shape[:2])
@@ -1224,10 +1181,8 @@ class DmdDenoisingStage(DenoisingStage):
progress_bar.update()
# Gather results if using sequence parallelism
if sp_group:
latents = sequence_model_parallel_all_gather(latents, dim=1)
latents = latents.permute(0, 2, 1, 3, 4)
# Update batch with final latents
batch.latents = latents
return batch
return batch
@@ -100,6 +100,92 @@ class ImageEncodingStage(PipelineStage):
return result
class Hy15ImageEncodingStage(ImageEncodingStage):
"""
Stage for encoding image prompts into embeddings for HunyuanVideo1.5 models.
"""
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
"""Verify image encoding stage inputs."""
return VerificationResult()
def forward(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> ForwardBatch:
"""
Encode the prompt into image encoder hidden states.
"""
if batch.pil_image is None:
batch.image_embeds = [
torch.zeros(1, 729, 1152, device=get_local_torch_device())
]
raw_latent_shape = list(batch.raw_latent_shape)
raw_latent_shape[1] = 1
batch.video_latent = torch.zeros(tuple(raw_latent_shape),
device=get_local_torch_device())
return batch
class MatrixGameImageEncodingStage(ImageEncodingStage):
CLIP_MEAN = [0.48145466, 0.4578275, 0.40821073]
CLIP_STD = [0.26862954, 0.26130258, 0.27577711]
@torch.no_grad()
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
assert batch.pil_image is not None
self.image_encoder = self.image_encoder.to(get_local_torch_device())
image = batch.pil_image
if isinstance(image, torch.Tensor):
if image.dim() == 5:
image = image[:, :,
0] # Extract first frame: [B, C, T, H, W] -> [B, C, H, W]
else:
from torchvision import transforms as T
transform = T.Compose([
T.ToTensor(), # [0, 1]
T.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5,
0.5]), # -> [-1, 1]
])
image = transform(image).unsqueeze(0) # [1, C, H, W]
device = get_local_torch_device()
image = image.to(device)
# F.interpolate with bicubic
image = torch.nn.functional.interpolate(image,
size=(224, 224),
mode='bicubic',
align_corners=False)
# [-1, 1] to [0, 1]
image = image * 0.5 + 0.5
# CLIP normalization
mean = torch.tensor(self.CLIP_MEAN, device=device,
dtype=image.dtype).view(1, 3, 1, 1)
std = torch.tensor(self.CLIP_STD, device=device,
dtype=image.dtype).view(1, 3, 1, 1)
image = (image - mean) / std
with set_forward_context(current_timestep=0, attn_metadata=None):
# WAN2_1ControlCLIPVisionConfig sets num_hidden_layers_override=31
# so last_hidden_state is the second-to-last layer output
outputs = self.image_encoder(pixel_values=image)
image_embeds = outputs.last_hidden_state
batch.image_embeds.append(image_embeds)
if fastvideo_args.image_encoder_cpu_offload:
self.image_encoder.to('cpu')
return batch
class RefImageEncodingStage(ImageEncodingStage):
"""
Stage for encoding reference image prompts into embeddings for Wan2.1 Control models.
@@ -526,3 +612,172 @@ class VideoVAEEncodingStage(ImageVAEEncodingStage):
result.add_check("video_latent", batch.video_latent,
[V.is_tensor, V.with_dims(5)])
return result
class MatrixGameImageVAEEncodingStage(ImageVAEEncodingStage):
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
assert batch.pil_image is not None
if fastvideo_args.mode == ExecutionMode.INFERENCE:
# Accept both PIL.Image and torch.Tensor
# Causal pipeline (is_causal=True) converts PIL to tensor in InputValidationStage
assert batch.pil_image is not None and isinstance(
batch.pil_image, PIL.Image.Image | torch.Tensor)
assert batch.height is not None and isinstance(batch.height, int)
assert batch.width is not None and isinstance(batch.width, int)
assert batch.num_frames is not None and isinstance(
batch.num_frames, int)
height = batch.height
width = batch.width
num_frames = batch.num_frames
elif fastvideo_args.mode == ExecutionMode.PREPROCESS:
assert batch.pil_image is not None and isinstance(
batch.pil_image, torch.Tensor)
assert batch.height is not None and isinstance(batch.height, list)
assert batch.width is not None and isinstance(batch.width, list)
assert batch.num_frames is not None and isinstance(
batch.num_frames, list)
num_frames = batch.num_frames[0]
height = batch.height[0]
width = batch.width[0]
else:
# Fallback for other modes
height = batch.height if isinstance(batch.height,
int) else batch.height[0]
width = batch.width if isinstance(batch.width,
int) else batch.width[0]
num_frames = batch.num_frames if isinstance(
batch.num_frames, int) else batch.num_frames[0]
self.vae = self.vae.to(get_local_torch_device())
# Process single image for I2V (latent dimensions computed but not used directly)
image = batch.pil_image
# Handle tensor input from causal pipeline
if isinstance(image, torch.Tensor):
# Causal pipeline provides tensor in [B, C, F, H, W] format
if image.dim() == 5:
# Already 5D, extract first frame for conditioning
# Shape: [B, C, F, H, W] -> use first frame [B, C, 1, H, W]
first_frame = image[:, :, :1] # Keep dim, [B, C, 1, H, W]
# Create video condition with first frame + zeros
video_condition = torch.cat([
first_frame,
first_frame.new_zeros(
first_frame.shape[0], first_frame.shape[1], num_frames -
1, first_frame.shape[3], first_frame.shape[4])
],
dim=2)
elif image.dim() == 4:
# [B, C, H, W] -> need to add frame dim
image = image.unsqueeze(2) # [B, C, 1, H, W]
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)
else:
raise ValueError(f"Unexpected tensor dimensions: {image.dim()}")
video_condition = video_condition.to(get_local_torch_device(),
dtype=torch.float32)
else:
# PIL Image input - use preprocess
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)
image = image.unsqueeze(2)
# Create video tensor with first frame as image, rest as zeros
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)
# 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 not vae_autocast_enabled:
video_condition = video_condition.to(vae_dtype)
encoder_output = self.vae.encode(video_condition)
# MatrixGame uses deterministic VAE encode for the first-frame conditioning.
# Sampling would inject random noise into the cond_concat tensor and destroy the action guidance.
img_cond = encoder_output.mode()
# manually using latents_mean and latents_std from config...
if (hasattr(self.vae.config, 'latents_mean')
and hasattr(self.vae.config, 'latents_std')):
# Convert config values to tensors
latents_mean = torch.tensor(self.vae.config.latents_mean,
device=img_cond.device,
dtype=img_cond.dtype).view(
1, -1, 1, 1, 1)
latents_std = torch.tensor(self.vae.config.latents_std,
device=img_cond.device,
dtype=img_cond.dtype).view(
1, -1, 1, 1, 1)
# Apply normalization: (latent - mean) * (1/std)
img_cond = (img_cond - latents_mean) / latents_std
elif (hasattr(self.vae, "shift_factor")
and self.vae.shift_factor is not None):
# Fallback to shift_factor/scaling_factor if available
if isinstance(self.vae.shift_factor, torch.Tensor):
img_cond -= self.vae.shift_factor.to(img_cond.device,
img_cond.dtype)
else:
img_cond -= self.vae.shift_factor
if hasattr(self.vae, 'scaling_factor'):
if isinstance(self.vae.scaling_factor, torch.Tensor):
img_cond = img_cond * self.vae.scaling_factor.to(
img_cond.device, img_cond.dtype)
else:
img_cond = img_cond * self.vae.scaling_factor
# Create mask_cond: ones for first frame, zeros for rest
# Shape: (B, 16, latent_frames, latent_height, latent_width)
mask_cond = torch.ones_like(img_cond)
mask_cond[:, :, 1:] = 0 # Set all frames except first to 0
# Create cond_concat: first 4 channels of mask + all 16 channels of img_cond
# Shape: (B, 20, latent_frames, latent_height, latent_width)
cond_concat = torch.cat([mask_cond[:, :4], img_cond], dim=1)
# Store cond_concat in batch.image_latent
# This will be concatenated with noise latents in DenoisingStage
batch.image_latent = cond_concat
# Offload models if needed
if hasattr(self, 'maybe_free_model_hooks'):
self.maybe_free_model_hooks()
self.vae.to("cpu")
return batch
+58 -12
View File
@@ -39,6 +39,7 @@ class InputValidationStage(PipelineStage):
assert seed is not None
seeds = [seed + i for i in range(num_videos_per_prompt)]
batch.seeds = seeds
# Peiyuan: using GPU seed will cause A100 and H100 to generate different results...
batch.generator = [
torch.Generator("cpu").manual_seed(seed) for seed in seeds
@@ -111,20 +112,28 @@ class InputValidationStage(PipelineStage):
) and batch.pil_image is not None:
img = batch.pil_image
ih, iw = img.height, img.width
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
max_area = 480 * 832
ow, oh = best_output_size(iw, ih, dw, dh, max_area)
scale = max(ow / iw, oh / ih)
img = img.resize((round(iw * scale), round(ih * scale)),
Image.LANCZOS)
pipeline_class_name = type(fastvideo_args.pipeline_config).__name__
if 'MatrixGame' in pipeline_class_name or 'MatrixCausal' in pipeline_class_name:
oh, ow = batch.height, batch.width
img = img.resize((ow, oh), Image.LANCZOS)
else:
# Standard Wan logic
patch_size = fastvideo_args.pipeline_config.dit_config.arch_config.patch_size
vae_stride = fastvideo_args.pipeline_config.vae_config.arch_config.scale_factor_spatial
dh, dw = patch_size[1] * vae_stride, patch_size[2] * vae_stride
max_area = 480 * 832
ow, oh = best_output_size(iw, ih, dw, dh, max_area)
scale = max(ow / iw, oh / ih)
img = img.resize((round(iw * scale), round(ih * scale)),
Image.LANCZOS)
# center-crop
x1 = (img.width - ow) // 2
y1 = (img.height - oh) // 2
img = img.crop((x1, y1, x1 + ow, y1 + oh))
# center-crop
x1 = (img.width - ow) // 2
y1 = (img.height - oh) // 2
img = img.crop((x1, y1, x1 + ow, y1 + oh))
assert img.width == ow and img.height == oh
logger.info("final processed img height: %s, img width: %s",
img.height, img.width)
@@ -186,6 +195,43 @@ class InputValidationStage(PipelineStage):
batch.video_latent = input_video
# Validate action control inputs (Matrix-Game)
if batch.mouse_cond is not None:
if batch.mouse_cond.dim() != 3 or batch.mouse_cond.shape[-1] != 2:
raise ValueError(
f"mouse_cond must have shape (B, T, 2), but got {batch.mouse_cond.shape}"
)
logger.info("Action control: mouse_cond validated - shape %s",
batch.mouse_cond.shape)
if batch.keyboard_cond is not None:
if batch.keyboard_cond.dim() != 3:
raise ValueError(
f"keyboard_cond must have 3 dimensions (B, T, K), but got {batch.keyboard_cond.dim()}"
)
keyboard_dim = batch.keyboard_cond.shape[-1]
if keyboard_dim not in {2, 4, 6, 7}:
raise ValueError(
f"keyboard_cond last dimension must be 2, 4, 6, or 7, but got {keyboard_dim}"
)
logger.info(
"Action control: keyboard_cond validated - shape %s (dim=%d)",
batch.keyboard_cond.shape, keyboard_dim)
if batch.grid_sizes is not None:
if not isinstance(batch.grid_sizes, list | tuple | torch.Tensor):
raise ValueError("grid_sizes must be a list, tuple, or tensor")
if isinstance(batch.grid_sizes, torch.Tensor):
if batch.grid_sizes.numel() != 3:
raise ValueError(
"grid_sizes must have 3 elements [F, H, W]")
else:
if len(batch.grid_sizes) != 3:
raise ValueError(
"grid_sizes must have 3 elements [F, H, W]")
logger.info("Action control: grid_sizes validated - %s",
batch.grid_sizes)
return batch
def verify_input(self, batch: ForwardBatch,
@@ -57,8 +57,17 @@ class LatentPreparationStage(PipelineStage):
# Adjust video length based on VAE version if needed
if hasattr(self, 'adjust_video_length'):
latent_num_frames = self.adjust_video_length(batch, fastvideo_args)
# Determine batch size
if isinstance(batch.prompt, list):
# Determine batch size; fall back to action/image inputs when no text encoder is present
if not batch.prompt_embeds:
if batch.keyboard_cond is not None:
batch_size = batch.keyboard_cond.shape[0]
elif batch.mouse_cond is not None:
batch_size = batch.mouse_cond.shape[0]
elif batch.image_embeds:
batch_size = batch.image_embeds[0].shape[0]
else:
batch_size = 1
elif isinstance(batch.prompt, list):
batch_size = len(batch.prompt)
elif batch.prompt is not None:
batch_size = 1
@@ -69,6 +78,19 @@ class LatentPreparationStage(PipelineStage):
batch_size *= batch.num_videos_per_prompt
# Get required parameters
if not batch.prompt_embeds:
# Create a dummy zero-length text embedding to satisfy downstream checks.
# Matrix-Game models have text_dim=0 and ignore encoder_hidden_states.
transformer_dtype = next(self.transformer.parameters()).dtype
device = get_local_torch_device()
dummy_prompt = torch.zeros(batch_size,
0,
self.transformer.hidden_size,
device=device,
dtype=transformer_dtype)
batch.prompt_embeds = [dummy_prompt]
batch.negative_prompt_embeds = []
batch.do_classifier_free_guidance = False
dtype = batch.prompt_embeds[0].dtype
device = get_local_torch_device()
generator = batch.generator
@@ -82,6 +104,7 @@ class LatentPreparationStage(PipelineStage):
raise ValueError("Height and width must be provided")
# Calculate latent shape
bcthw_shape: tuple[int, ...] | None = None
if self.use_btchw_layout:
shape = (
batch_size,
@@ -92,6 +115,7 @@ class LatentPreparationStage(PipelineStage):
width // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
)
bcthw_shape = tuple(shape[i] for i in [0, 2, 1, 3, 4])
else:
shape = (
batch_size,
@@ -102,6 +126,7 @@ class LatentPreparationStage(PipelineStage):
width // fastvideo_args.pipeline_config.vae_config.arch_config.
spatial_compression_ratio,
)
bcthw_shape = shape
# Validate generator if it's a list
if isinstance(generator, list) and len(generator) != batch_size:
@@ -123,7 +148,7 @@ class LatentPreparationStage(PipelineStage):
latents = latents * self.scheduler.init_noise_sigma
# Update batch with prepared latents
batch.latents = latents
batch.raw_latent_shape = latents.shape
batch.raw_latent_shape = bcthw_shape
return batch
@@ -154,10 +179,12 @@ class LatentPreparationStage(PipelineStage):
"""Verify latent preparation stage inputs."""
result = VerificationResult()
result.add_check(
"prompt_or_embeds", None, lambda _: V.string_or_list_strings(
batch.prompt) or V.list_not_empty(batch.prompt_embeds))
result.add_check("prompt_embeds", batch.prompt_embeds,
V.list_of_tensors)
"prompt_or_embeds", None,
lambda _: V.string_or_list_strings(batch.prompt) or not batch.
prompt_embeds or V.list_not_empty(batch.prompt_embeds))
if batch.prompt_embeds:
result.add_check("prompt_embeds", batch.prompt_embeds,
V.list_of_tensors)
result.add_check("num_videos_per_prompt", batch.num_videos_per_prompt,
V.positive_int)
result.add_check("generator", batch.generator,
@@ -0,0 +1,566 @@
from __future__ import annotations
from typing import Any
import torch # type: ignore
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.utils import pred_noise_to_pred_video, pred_noise_to_x_bound
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 (
SlidingTileAttentionBackend)
st_attn_available = True
except ImportError:
st_attn_available = False
SlidingTileAttentionBackend = None # type: ignore
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionBackend)
vsa_available = True
except ImportError:
vsa_available = False
VideoSparseAttentionBackend = None # type: ignore
logger = init_logger(__name__)
class MatrixGameCausalDenoisingStage(DenoisingStage):
def __init__(self,
transformer,
scheduler,
pipeline=None,
transformer_2=None,
vae=None) -> None:
super().__init__(transformer, scheduler, pipeline, transformer_2, vae)
self.transformer = transformer
self.transformer_2 = transformer_2
self.vae = vae
self.num_transformer_blocks = len(self.transformer.blocks)
if hasattr(self.transformer, 'config') and hasattr(
self.transformer.config, 'arch_config'):
self.num_frame_per_block = getattr(
self.transformer.config.arch_config, 'num_frames_per_block',
getattr(self.transformer, 'num_frame_per_block', 1))
self.sliding_window_num_frames = getattr(
self.transformer.config.arch_config,
'sliding_window_num_frames', 15)
else:
self.num_frame_per_block = getattr(self.transformer,
'num_frame_per_block', 1)
self.sliding_window_num_frames = 15
try:
self.local_attn_size = getattr(self.transformer, "local_attn_size",
-1)
except Exception:
self.local_attn_size = -1
assert self.local_attn_size != -1, (
f"local_attn_size must be set for Matrix-Game causal inference, "
f"got {self.local_attn_size}. Check MatrixGameWanVideoArchConfig.")
assert self.num_frame_per_block > 0, (
f"num_frame_per_block must be positive, got {self.num_frame_per_block}"
)
logger.info(
"MatrixGame causal inference initialized: "
"local_attn_size=%s, num_frame_per_block=%s", self.local_attn_size,
self.num_frame_per_block)
self.action_config = getattr(self.transformer, 'action_config', {})
self.use_action_module = len(self.action_config) > 0
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
target_dtype = torch.bfloat16
autocast_enabled = (target_dtype != torch.float32
) and not fastvideo_args.disable_autocast
latent_seq_length = batch.latents.shape[-1] * batch.latents.shape[-2]
patch_size = self.transformer.patch_size
patch_ratio = patch_size[-1] * patch_size[-2]
self.frame_seq_length = latent_seq_length // patch_ratio
independent_first_frame = getattr(self.transformer,
'independent_first_frame', False)
timesteps = torch.tensor(
fastvideo_args.pipeline_config.dmd_denoising_steps,
dtype=torch.long).cpu()
if fastvideo_args.pipeline_config.warp_denoising_step:
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(),
torch.tensor([0],
dtype=torch.float32)))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(get_local_torch_device())
boundary_ratio = getattr(fastvideo_args.pipeline_config.dit_config,
'boundary_ratio', None)
if boundary_ratio is not None:
boundary_timestep = boundary_ratio * self.scheduler.num_train_timesteps
high_noise_timesteps = timesteps[timesteps >= boundary_timestep]
else:
boundary_timestep = None
high_noise_timesteps = None
image_embeds = batch.image_embeds
if len(image_embeds) > 0:
assert torch.isnan(image_embeds[0]).sum() == 0
image_embeds = [
image_embed.to(target_dtype) for image_embed in image_embeds
]
# directly set the kwarg.
image_kwargs = {"encoder_hidden_states_image": image_embeds}
pos_cond_kwargs: dict[str, Any] = {}
if st_attn_available and self.attn_backend == SlidingTileAttentionBackend:
self.prepare_sta_param(batch, fastvideo_args)
assert batch.latents is not None, "latents must be provided"
latents = batch.latents
b, c, t, h, w = latents.shape
prompt_embeds = batch.prompt_embeds
assert torch.isnan(prompt_embeds[0]).sum() == 0
kv_cache1 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache2 = None
if boundary_timestep is not None:
kv_cache2 = self._initialize_kv_cache(batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
kv_cache_mouse = None
kv_cache_keyboard = None
if self.use_action_module:
kv_cache_mouse, kv_cache_keyboard = self._initialize_action_kv_cache(
batch_size=latents.shape[0],
dtype=target_dtype,
device=latents.device)
def _get_kv_cache(timestep_val: float) -> list[dict]:
if boundary_timestep is not None:
if timestep_val >= boundary_timestep:
return kv_cache1
else:
assert kv_cache2 is not None, "kv_cache2 is not initialized"
return kv_cache2
return kv_cache1
crossattn_cache = self._initialize_crossattn_cache(
batch_size=latents.shape[0],
max_text_len=257, # 1 CLS + 256 patch tokens
dtype=target_dtype,
device=latents.device)
pos_start_base = 0
if t % self.num_frame_per_block != 0:
raise ValueError(
"num_frames must be divisible by num_frame_per_block for causal denoising"
)
num_blocks = t // self.num_frame_per_block
block_sizes = [self.num_frame_per_block] * num_blocks
start_index = 0
if boundary_timestep is not None:
block_sizes[0] = 1
# NOTE: MatrixGame does NOT process the first frame separately.
# The first frame information is already encoded in batch.image_latent (cond_concat)
# and will be used by the model via channel concatenation: torch.cat([x, cond_concat], dim=1)
with self.progress_bar(total=len(block_sizes) *
len(timesteps)) as progress_bar:
for block_idx, current_num_frames in enumerate(block_sizes):
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
_video_raw_latent_shape = noise_latents_btchw.shape # noqa: F841
# NOTE: crossattn_cache should NOT be reset between blocks!
action_kwargs = self._prepare_action_kwargs(
batch, start_index, current_num_frames)
for i, t_cur in enumerate(timesteps):
if boundary_timestep is not None and t_cur < boundary_timestep:
current_model = self.transformer_2 if self.transformer_2 is not None else self.transformer
else:
current_model = self.transformer
noise_latents = noise_latents_btchw.clone()
latent_model_input = current_latents.to(target_dtype)
if batch.image_latent is not None and independent_first_frame and start_index == 0:
latent_model_input = torch.cat([
latent_model_input,
batch.image_latent.to(target_dtype)
],
dim=2)
# t_expand needs to be [batch * frames] to match flattened pred_noise/noise_latents
t_expand = t_cur.repeat(latent_model_input.shape[0] *
current_num_frames)
if vsa_available and self.attn_backend == VideoSparseAttentionBackend:
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(
)
attn_metadata = self.attn_metadata_builder.build(
current_timestep=i,
raw_latent_shape=(current_num_frames, h, w),
patch_size=fastvideo_args.pipeline_config.
dit_config.patch_size,
STA_param=batch.STA_param,
VSA_sparsity=fastvideo_args.VSA_sparsity,
device=get_local_torch_device(),
)
assert attn_metadata is not None, "attn_metadata cannot be None"
else:
attn_metadata = None
else:
attn_metadata = None
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=i,
attn_metadata=attn_metadata,
forward_batch=batch):
# Expand timestep to per-frame format [batch, num_frames] for causal model
t_expanded_noise = t_cur * torch.ones(
(latent_model_input.shape[0], current_num_frames),
device=latent_model_input.device,
dtype=torch.long)
model_kwargs = {
"kv_cache": _get_kv_cache(t_cur),
"crossattn_cache": crossattn_cache,
"current_start": (pos_start_base + start_index) *
self.frame_seq_length,
"start_frame": start_index,
}
if self.use_action_module and current_model == self.transformer:
model_kwargs.update({
"kv_cache_mouse":
kv_cache_mouse,
"kv_cache_keyboard":
kv_cache_keyboard,
})
model_kwargs.update(action_kwargs)
pred_noise_btchw = current_model(
latent_model_input,
prompt_embeds,
t_expanded_noise,
**image_kwargs,
**pos_cond_kwargs,
**model_kwargs,
).permute(0, 2, 1, 3, 4)
if boundary_timestep is not None and t_cur >= boundary_timestep:
pred_video_btchw = pred_noise_to_x_bound(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
boundary_timestep=torch.ones_like(t_expand) *
boundary_timestep,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
else:
pred_video_btchw = pred_noise_to_pred_video(
pred_noise=pred_noise_btchw.flatten(0, 1),
noise_input_latent=noise_latents.flatten(0, 1),
timestep=t_expand,
scheduler=self.scheduler).unflatten(
0, pred_noise_btchw.shape[:2])
if i < len(timesteps) - 1:
next_timestep = timesteps[i + 1] * torch.ones(
[1],
dtype=torch.long,
device=pred_video_btchw.device)
noise = torch.randn(
pred_video_btchw.shape,
dtype=pred_video_btchw.dtype,
generator=(batch.generator[0] if isinstance(
batch.generator, list) else
batch.generator)).to(
pred_video_btchw.device)
noise_btchw = noise
if boundary_timestep is not None and high_noise_timesteps is not None and i < len(
high_noise_timesteps) - 1:
noise_latents_btchw = self.scheduler.add_noise_high(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1), next_timestep,
torch.ones_like(next_timestep) *
boundary_timestep).unflatten(
0, pred_video_btchw.shape[:2])
elif boundary_timestep is not None and high_noise_timesteps is not None and i == len(
high_noise_timesteps) - 1:
noise_latents_btchw = pred_video_btchw
else:
noise_latents_btchw = self.scheduler.add_noise(
pred_video_btchw.flatten(0, 1),
noise_btchw.flatten(0, 1),
next_timestep).unflatten(
0, pred_video_btchw.shape[:2])
current_latents = noise_latents_btchw.permute(
0, 2, 1, 3, 4)
else:
current_latents = pred_video_btchw.permute(
0, 2, 1, 3, 4)
if progress_bar is not None:
progress_bar.update()
latents[:, :, start_index:start_index +
current_num_frames, :, :] = current_latents
context_noise = getattr(fastvideo_args.pipeline_config,
"context_noise", 0)
# Expand context timestep to per-frame format [batch, num_frames] for causal model
t_context = torch.ones([latents.shape[0], current_num_frames],
device=latents.device,
dtype=torch.long) * int(context_noise)
context_bcthw = current_latents.to(target_dtype)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
set_forward_context(current_timestep=0,
attn_metadata=attn_metadata,
forward_batch=batch):
context_model_kwargs = {
"kv_cache": kv_cache1,
"crossattn_cache": crossattn_cache,
"current_start":
(pos_start_base + start_index) * self.frame_seq_length,
"start_frame": start_index,
}
if self.use_action_module:
context_model_kwargs.update({
"kv_cache_mouse":
kv_cache_mouse,
"kv_cache_keyboard":
kv_cache_keyboard,
})
context_model_kwargs.update(action_kwargs)
if boundary_timestep is not None and self.transformer_2 is not None:
self.transformer_2(
context_bcthw,
prompt_embeds,
t_context,
kv_cache=kv_cache2,
crossattn_cache=crossattn_cache,
current_start=(pos_start_base + start_index) *
self.frame_seq_length,
start_frame=start_index,
**image_kwargs,
**pos_cond_kwargs,
)
self.transformer(
context_bcthw,
prompt_embeds,
t_context,
**image_kwargs,
**pos_cond_kwargs,
**context_model_kwargs,
)
start_index += current_num_frames
if boundary_timestep is not None:
num_frames_to_remove = self.num_frame_per_block - 1
if num_frames_to_remove > 0:
latents = latents[:, :, :-num_frames_to_remove, :, :]
batch.latents = latents
return batch
def _prepare_action_kwargs(self, batch: ForwardBatch, start_index: int,
num_frames: int) -> dict[str, Any]:
action_kwargs: dict[str, Any] = {}
if not self.use_action_module:
return action_kwargs
vae_time_compression_ratio = 4
end_frame_idx = 1 + vae_time_compression_ratio * (start_index +
num_frames - 1)
if hasattr(batch, 'mouse_cond') and batch.mouse_cond is not None:
action_kwargs['mouse_cond'] = batch.mouse_cond[:, :end_frame_idx]
if hasattr(batch, 'keyboard_cond') and batch.keyboard_cond is not None:
action_kwargs[
'keyboard_cond'] = batch.keyboard_cond[:, :end_frame_idx]
# CRITICAL: Pass num_frame_per_block to model - this should be num_frames (current block size)
action_kwargs['num_frame_per_block'] = num_frames
return action_kwargs
def _initialize_kv_cache(self, batch_size: int, dtype: torch.dtype,
device: torch.device) -> list[dict]:
kv_cache = []
num_attention_heads = self.transformer.num_attention_heads
attention_head_dim = getattr(
self.transformer, 'attention_head_dim',
self.transformer.hidden_size // num_attention_heads)
if self.local_attn_size != -1:
kv_cache_size = self.local_attn_size * self.frame_seq_length
else:
kv_cache_size = self.frame_seq_length * self.sliding_window_num_frames
for _ in range(self.num_transformer_blocks):
kv_cache.append({
"k":
torch.zeros([
batch_size, kv_cache_size, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"v":
torch.zeros([
batch_size, kv_cache_size, num_attention_heads,
attention_head_dim
],
dtype=dtype,
device=device),
"global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
"local_end_index":
torch.tensor([0], dtype=torch.long, device=device),
})
return kv_cache
def _initialize_action_kv_cache(self, batch_size: int, dtype: torch.dtype,
device: torch.device):
kv_cache_mouse = []
kv_cache_keyboard = []
action_heads = self.action_config.get('heads_num', 16)
mouse_head_dim = self.action_config.get('mouse_hidden_dim',
1024) // action_heads
keyboard_head_dim = self.action_config.get('keyboard_hidden_dim',
1024) // action_heads
if self.local_attn_size != -1:
kv_cache_size = self.local_attn_size
else:
kv_cache_size = 15
for _ in range(self.num_transformer_blocks):
kv_cache_keyboard.append({
"k":
torch.zeros([
batch_size, kv_cache_size, action_heads, keyboard_head_dim
],
dtype=dtype,
device=device),
"v":
torch.zeros([
batch_size, kv_cache_size, action_heads, keyboard_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),
})
kv_cache_mouse.append({
"k":
torch.zeros([
batch_size * self.frame_seq_length, kv_cache_size,
action_heads, mouse_head_dim
],
dtype=dtype,
device=device),
"v":
torch.zeros([
batch_size * self.frame_seq_length, kv_cache_size,
action_heads, mouse_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_cache_mouse, kv_cache_keyboard
def _initialize_crossattn_cache(self, batch_size: int, max_text_len: int,
dtype: torch.dtype,
device: torch.device) -> list[dict]:
crossattn_cache = []
num_attention_heads = self.transformer.num_attention_heads
attention_head_dim = getattr(
self.transformer, 'attention_head_dim',
self.transformer.hidden_size // num_attention_heads)
for _ in range(self.num_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
def verify_input(self, batch: ForwardBatch,
fastvideo_args: FastVideoArgs) -> VerificationResult:
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
+52 -10
View File
@@ -6,6 +6,7 @@ This module contains implementations of prompt encoding stages for diffusion pip
"""
import torch
from typing import Any
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
@@ -100,9 +101,9 @@ class TextEncodingStage(PipelineStage):
"""Verify text encoding stage inputs."""
result = VerificationResult()
result.add_check("prompt", batch.prompt, V.string_or_list_strings)
result.add_check(
"negative_prompt", batch.negative_prompt, lambda x: not batch.
do_classifier_free_guidance or V.string_not_empty(x))
# result.add_check(
# "negative_prompt", batch.negative_prompt, lambda x: not batch.
# do_classifier_free_guidance or V.string_not_empty(x))
result.add_check("do_classifier_free_guidance",
batch.do_classifier_free_guidance, V.bool_value)
result.add_check("prompt_embeds", batch.prompt_embeds, V.is_list)
@@ -203,20 +204,45 @@ class TextEncodingStage(PipelineStage):
preprocess_func = preprocess_funcs[i]
postprocess_func = postprocess_funcs[i]
processed_texts: list[str] = []
for prompt_str in texts:
processed_texts.append(preprocess_func(prompt_str))
tok_kwargs = dict(encoder_config.tokenizer_kwargs)
if max_length is not None:
tok_kwargs["max_length"] = max_length
elif hasattr(fastvideo_args.pipeline_config,
"text_encoder_max_lengths"):
tok_kwargs[
"max_length"] = fastvideo_args.pipeline_config.text_encoder_max_lengths[
i]
if truncation is not None:
tok_kwargs["truncation"] = truncation
if padding is not None:
tok_kwargs["padding"] = padding
text_inputs = tokenizer(processed_texts,
**tok_kwargs).to(target_device)
processed_texts: list[str] = []
for prompt_str in texts:
processed_text = preprocess_func(prompt_str)
if processed_text is not None:
processed_texts.append(processed_text)
else:
# Assuming batch_size = 1
prompt_embeds = torch.zeros((1, tok_kwargs["max_length"],
encoder_config.hidden_size),
device=target_device)
attention_mask = torch.zeros((1, tok_kwargs["max_length"]),
device=target_device,
dtype=torch.int64)
embeds_list.append(prompt_embeds)
attn_masks_list.append(attention_mask)
return self.return_embeds(embeds_list, attn_masks_list,
return_type,
return_attention_mask, indices)
if encoder_config.is_chat_model:
text_inputs = tokenizer.apply_chat_template(
processed_texts, **tok_kwargs).to(target_device)
else:
text_inputs = tokenizer(processed_texts,
**tok_kwargs).to(target_device)
input_ids = text_inputs["input_ids"]
attention_mask = text_inputs["attention_mask"]
@@ -228,13 +254,29 @@ class TextEncodingStage(PipelineStage):
output_hidden_states=True,
)
prompt_embeds = postprocess_func(outputs)
try:
prompt_embeds = postprocess_func(outputs)
except Exception:
prompt_embeds, attention_mask = postprocess_func(
outputs, attention_mask)
if dtype is not None:
prompt_embeds = prompt_embeds.to(dtype=dtype)
embeds_list.append(prompt_embeds)
if return_attention_mask:
attn_masks_list.append(attention_mask)
return self.return_embeds(embeds_list, attn_masks_list, return_type,
return_attention_mask, indices)
def return_embeds(
self,
embeds_list: list[torch.Tensor],
attn_masks_list: list[torch.Tensor],
return_type: str = "list",
return_attention_mask: bool = False,
indices: list[int] | None = None,
) -> Any:
# Shape results according to return_type
if return_type == "list":
if return_attention_mask:
+1 -1
View File
@@ -183,7 +183,7 @@ class CudaPlatformBase(Platform):
raise ImportError(
"The Video Sparse Attention backend is not installed. "
"To install it, please follow the instructions at: "
"https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation.html "
"https://hao-ai-lab.github.io/FastVideo/video_sparse_attention/installation "
) from e
elif selected_backend == AttentionBackendEnum.VMOBA_ATTN:
+5
View File
@@ -208,6 +208,11 @@ def get_or_create_profiler(trace_dir: str | None) -> TorchProfilerController:
profile_memory=envs.FASTVIDEO_TORCH_PROFILER_WITH_PROFILE_MEMORY,
with_stack=envs.FASTVIDEO_TORCH_PROFILER_WITH_STACK,
with_flops=envs.FASTVIDEO_TORCH_PROFILER_WITH_FLOPS,
schedule=torch.profiler.schedule(
wait=envs.FASTVIDEO_TORCH_PROFILER_WAIT_STEPS,
warmup=envs.FASTVIDEO_TORCH_PROFILER_WARMUP_STEPS,
active=envs.FASTVIDEO_TORCH_PROFILER_ACTIVE_STEPS,
),
on_trace_ready=torch.profiler.tensorboard_trace_handler(trace_dir,
use_gzip=True),
)
@@ -0,0 +1,140 @@
# SPDX-License-Identifier: Apache-2.0
import os
import numpy as np
import pytest
import torch
from torch.distributed.tensor import DTensor
from torch.testing import assert_close
from transformers import AutoConfig, AutoTokenizer, T5EncoderModel
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TextEncoderLoader
from fastvideo.utils import maybe_download_model, PRECISION_TO_TYPE
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.configs.models.encoders import T5Config
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29503"
@pytest.fixture
def t5_model_paths():
base_model_path = "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
model_path = maybe_download_model(base_model_path)
text_encoder_path = os.path.join(model_path, "text_encoder_2")
tokenizer_path = os.path.join(model_path, "tokenizer_2")
return text_encoder_path, tokenizer_path
@pytest.mark.usefixtures("distributed_setup")
def test_t5_encoder(t5_model_paths):
# Initialize the two model implementations
text_encoder_path, tokenizer_path = t5_model_paths
hf_config = AutoConfig.from_pretrained(text_encoder_path)
print(hf_config)
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
precision_str = "fp32"
precision = PRECISION_TO_TYPE[precision_str]
model1 = T5EncoderModel.from_pretrained(text_encoder_path).to(
precision).to(device).eval()
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
args = FastVideoArgs(model_path=text_encoder_path,
pipeline_config=PipelineConfig(text_encoder_configs=(T5Config(),),
text_encoder_precisions=(precision_str,)),
pin_cpu_memory=False)
loader = TextEncoderLoader()
model2 = loader.load(text_encoder_path, args)
model2 = model2.to(precision)
model2.eval()
# Sanity check weights between the two models
logger.info("Comparing model weights for sanity check...")
params1 = dict(model1.named_parameters())
params2 = dict(model2.named_parameters())
# Check number of parameters
logger.info("Model1 has %s parameters", len(params1))
logger.info("Model2 has %s parameters", len(params2))
# check if embed_tokens are the same
weights = ["encoder.block.{}.layer.0.layer_norm.weight", \
"encoder.block.{}.layer.0.SelfAttention.o.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_0.weight", "encoder.block.{}.layer.1.DenseReluDense.wi_1.weight",\
"encoder.block.{}.layer.1.DenseReluDense.wo.weight", \
"encoder.block.{}.layer.1.layer_norm.weight", "encoder.final_layer_norm.weight"]
for idx in range(hf_config.num_hidden_layers):
for w in weights:
name1 = w.format(idx)
name2 = w.format(idx)
p1 = params1[name1]
p2 = params2[name2]
p2 = (p2.to_local() if isinstance(p2, DTensor) else p2).to(p1)
assert_close(p1, p2, atol=1e-4, rtol=1e-4)
# Test with some sample prompts
prompts = [
"Once upon a time", "The quick brown fox jumps over",
"In a galaxy far, far away"
]
logger.info("Testing T5 encoder with sample prompts")
with torch.no_grad():
for prompt in prompts:
logger.info("Testing prompt: %s", prompt)
# Tokenize the prompt
tokens = tokenizer(prompt,
padding="max_length",
max_length=512,
truncation=True,
add_special_tokens=True,
return_tensors="pt").to(device)
# Get outputs from HuggingFace implementation
# filter out padding input_ids
# tokens.input_ids = tokens.input_ids[tokens.attention_mask==1]
# tokens.attention_mask = tokens.attention_mask[tokens.attention_mask==1]
outputs1 = model1(input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask.float())[0]
print("--------------------------------")
logger.info("Testing model2")
# Get outputs from our implementation
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs2 = model2(
input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask,
).last_hidden_state
# Compare last hidden states
last_hidden_state1 = outputs1[tokens.attention_mask == 1]
last_hidden_state2 = outputs2[tokens.attention_mask == 1]
assert last_hidden_state1.shape == last_hidden_state2.shape, \
f"Hidden state shapes don't match: {last_hidden_state1.shape} vs {last_hidden_state2.shape}"
max_diff_hidden = torch.max(
torch.abs(last_hidden_state1 - last_hidden_state2))
mean_diff_hidden = torch.mean(
torch.abs(last_hidden_state1 - last_hidden_state2))
logger.info("Maximum difference in last hidden states: %s",
max_diff_hidden.item())
logger.info("Mean difference in last hidden states: %s",
mean_diff_hidden.item())
logger.info("Max memory allocated: %s GB", torch.cuda.max_memory_allocated() / 1024**3)
# Check if outputs are similar (allowing for small numerical differences)
assert mean_diff_hidden < 1e-4, \
f"Hidden states differ significantly: mean diff = {mean_diff_hidden.item()}"
assert max_diff_hidden < 1e-4, \
f"Hidden states differ significantly: max diff = {max_diff_hidden.item()}"
@@ -0,0 +1,150 @@
# SPDX-License-Identifier: Apache-2.0
import os
import pytest
import torch
from torch.distributed.tensor import DTensor
from torch.testing import assert_close
from transformers import AutoConfig, AutoTokenizer, Qwen2_5_VLTextModel
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
from fastvideo.models.loader.component_loader import TextEncoderLoader
from fastvideo.utils import maybe_download_model, PRECISION_TO_TYPE
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.configs.models.encoders import Qwen2_5_VLConfig
logger = init_logger(__name__)
os.environ["MASTER_ADDR"] = "localhost"
os.environ["MASTER_PORT"] = "29505"
@pytest.fixture
def qwen_model_path():
base_model_path = "hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v"
model_path = maybe_download_model(base_model_path)
text_encoder_path = os.path.join(model_path, "text_encoder")
tokenizer_path = os.path.join(model_path, "tokenizer")
return text_encoder_path, tokenizer_path
@pytest.mark.usefixtures("distributed_setup")
def test_qwen2_5_encoder(qwen_model_path):
text_encoder_path, tokenizer_path = qwen_model_path
hf_config = AutoConfig.from_pretrained(text_encoder_path)
print(hf_config)
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
# Qwen2.5-VL default dtype is usually bf16
precision_str = "fp32"
precision = PRECISION_TO_TYPE[precision_str]
logger.info(f"Using precision: {precision_str}")
# Load HF model (Base model)
model1 = Qwen2_5_VLTextModel.from_pretrained(text_encoder_path).to(
precision).to(device).eval()
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
# Load FastVideo model
args = FastVideoArgs(model_path=text_encoder_path,
pipeline_config=PipelineConfig(text_encoder_configs=(Qwen2_5_VLConfig(),),
text_encoder_precisions=(precision_str,)),
pin_cpu_memory=False)
loader = TextEncoderLoader()
model2 = loader.load(text_encoder_path, args)
model2 = model2.to(precision)
model2.eval()
# Sanity check weights
logger.info("Comparing model weights for sanity check...")
params1 = dict(model1.named_parameters())
params2 = dict(model2.named_parameters())
logger.info("Model1 has %s parameters", len(params1))
logger.info("Model2 has %s parameters", len(params2))
# Check common layers like Norms which are likely not merged/sharded in a way that changes name significantly
# or simple linear layers if names match.
# Note: FastVideo uses QKVParallelLinear, so q_proj, k_proj, v_proj are merged.
# HF Qwen2_5_VL uses separate projections? No, usually they are separate nn.Linear in HF.
weights_to_check = [
"norm.weight",
"layers.{}.self_attn.o_proj.weight",
"layers.{}.input_layernorm.weight",
"layers.{}.post_attention_layernorm.weight",
"layers.{}.mlp.down_proj.weight"
]
for idx in range(hf_config.num_hidden_layers):
for w in weights_to_check:
name1 = w.format(idx)
name2 = w.format(idx)
p1 = params1[name1]
p2 = params2[name2]
p2 = (p2.to_local() if isinstance(p2, DTensor) else p2).to(p1)
# Check shape
assert p1.shape == p2.shape, f"Shape mismatch for {w}: {p1.shape} vs {p2.shape}"
# Check values
assert_close(p1, p2, atol=1e-7, rtol=1e-7, msg=f"Weight mismatch for {w}")
# Test with sample prompts
prompts = [
"Hello world",
"The quick brown fox jumps over the lazy dog."
]
logger.info("Testing with sample prompts")
with torch.no_grad():
for prompt in prompts:
logger.info(f"Prompt: {prompt}")
tokens = tokenizer(prompt, return_tensors="pt", padding="max_length", max_length=1000, truncation=True).to(device)
# HF Forward
# AutoModel for Qwen2.5-VL usually returns BaseModelOutputWithPast
outputs1 = model1(
input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask,
output_hidden_states=True
).hidden_states[-3]
# FastVideo Forward
with set_forward_context(current_timestep=0, attn_metadata=None):
outputs2 = model2(
input_ids=tokens.input_ids,
attention_mask=tokens.attention_mask,
output_hidden_states=True
).hidden_states[-3]
# Compare
# Filter padding for comparison if needed, but here we just check raw output matching
# Check shapes
assert outputs1.shape == outputs2.shape, f"Output shape mismatch: {outputs1.shape} vs {outputs2.shape}"
diff = torch.abs(outputs1 - outputs2)
max_diff = diff.max().item()
mean_diff = diff.mean().item()
logger.info(f"Max diff: {max_diff}")
logger.info(f"Mean diff: {mean_diff}")
# Thresholds
# Qwen2.5-VL RoPE is complex, if our implementation is slightly off (e.g. float32 conversion logic in RoPE),
# differences might appear. But should be small.
if precision_str == "bf16":
atol = 5e-2 # relaxed for bf16
else:
atol = 1e-3
if max_diff > atol:
logger.warning(f"Max diff {max_diff} > {atol}. Checking if it's acceptable...")
# If mean diff is small, maybe just outliers
assert mean_diff < atol, f"Mean diff {mean_diff} too high"
else:
logger.info("Outputs match within tolerance.")
@@ -27,7 +27,7 @@ def test_inference_vmoba():
"--height", "480",
"--width", "832",
"--num-frames", "77",
"--num-inference-steps", "50",
"--num-inference-steps", "10",
"--moba-config-path", moba_config,
"--fps", "16",
"--guidance-scale", "6.0",
@@ -124,7 +124,7 @@ def run_training():
"--num_height", "480",
"--num_width", "832",
"--num_frames", "81",
"--validation_guidance_scale", "1.0",
"--validation_guidance_scale", "6.0",
"--num_euler_timesteps", "50",
"--multi_phased_distill_schedule", "4000-1",
"--weight_decay", "0.01",
@@ -140,8 +140,8 @@ def run_training():
def test_e2e_overfit_single_sample():
os.environ["WANDB_MODE"] = "online"
# download_data()
# run_preprocessing()
download_data()
run_preprocessing()
run_training()
reference_video_file = os.path.join(os.path.dirname(__file__), "reference_video_1_sample_v0.mp4")
@@ -114,7 +114,7 @@ def run_training():
"--dataloader_num_workers", "10",
"--gradient_accumulation_steps", "1",
"--max_train_steps", "901",
"--learning_rate", "1e-5",
"--learning_rate", "5e-6",
"--mixed_precision", "bf16",
"--weight_only_checkpointing_steps", "6000",
"--training_state_checkpointing_steps", "6000",
@@ -129,7 +129,7 @@ def run_training():
"--num_height", "480",
"--num_width", "832",
"--num_frames", "81",
"--validation_guidance_scale", "1.0",
"--validation_guidance_scale", "3.0",
"--num_euler_timesteps", "50",
"--multi_phased_distill_schedule", "4000-1",
"--weight_decay", "0.01",
@@ -21,7 +21,8 @@ elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
else:
# device_reference_folder = "L40S" + device_reference_folder_suffix
raise ValueError(f"Unsupported device for ssim tests: {device_name}")
logger.warning(f"Unsupported device for ssim tests: {device_name}")
# raise ValueError(f"Unsupported device for ssim tests: {device_name}")
# Base parameters from the shell script
@@ -21,7 +21,8 @@ elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
else:
# device_reference_folder = "L40S" + device_reference_folder_suffix
raise ValueError(f"Unsupported device for ssim tests: {device_name}")
logger.warning(f"Unsupported device for ssim tests: {device_name}")
# raise ValueError(f"Unsupported device for ssim tests: {device_name}")
# Base parameters from the shell script
HUNYUAN_PARAMS = {
@@ -0,0 +1,176 @@
# SPDX-License-Identifier: Apache-2.0
import os
import torch
import pytest
from fastvideo import VideoGenerator
from fastvideo.models.dits.matrix_game.utils import create_action_presets
from fastvideo.logger import init_logger
from fastvideo.tests.utils import (
compute_video_ssim_torchvision,
write_ssim_results,
)
from fastvideo.worker.multiproc_executor import MultiprocExecutor
logger = init_logger(__name__)
device_name = torch.cuda.get_device_name()
device_reference_folder_suffix = "_reference_videos"
if "A40" in device_name:
device_reference_folder = "A40" + device_reference_folder_suffix
elif "L40S" in device_name:
device_reference_folder = "L40S" + device_reference_folder_suffix
else:
# device_reference_folder = "L40S" + device_reference_folder_suffix
logger.warning(f"Unsupported device for ssim tests: {device_name}")
# raise ValueError(f"Unsupported device for ssim tests: {device_name}")
# Base parameters from the shell script
MATRIXGAME_PARAMS = {
"num_gpus": 1,
"model_path": "FastVideo/Matrix-Game-2.0-Base-Diffusers",
"height": 352,
"width": 640,
"num_frames": 117,
"num_inference_steps": 10,
"seed": 1024,
"keyboard_dim": 4,
}
MODEL_TO_PARAMS = {
"Matrix-Game-2.0-Diffusers-Base": MATRIXGAME_PARAMS,
}
I2V_MODEL_TO_PARAMS = {}
TEST_PROMPTS = [
"", # MatrixGame is I2V with action conditions, no text prompt
]
TEST_IMAGE_PATHS = [
"https://raw.githubusercontent.com/SkyworkAI/Matrix-Game/main/Matrix-Game-2/demo_images/universal/0002.png",
]
@pytest.mark.parametrize("prompt", TEST_PROMPTS)
@pytest.mark.parametrize("ATTENTION_BACKEND", ["FLASH_ATTN"])
@pytest.mark.parametrize("model_id", list(MODEL_TO_PARAMS.keys()))
def test_matrixgame_similarity(prompt, ATTENTION_BACKEND, model_id):
"""
Test that runs inference with different parameters and compares the output
to reference videos using SSIM.
"""
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = ATTENTION_BACKEND
script_dir = os.path.dirname(os.path.abspath(__file__))
base_output_dir = os.path.join(script_dir, "generated_videos", model_id)
output_dir = os.path.join(base_output_dir, ATTENTION_BACKEND)
output_video_name = "video.mp4"
os.makedirs(output_dir, exist_ok=True)
BASE_PARAMS = MODEL_TO_PARAMS[model_id]
num_inference_steps = BASE_PARAMS["num_inference_steps"]
# Create action conditions for MatrixGame
actions = create_action_presets(
BASE_PARAMS["num_frames"], keyboard_dim=BASE_PARAMS["keyboard_dim"],
seed=BASE_PARAMS["seed"]
)
latent_frames = (BASE_PARAMS["num_frames"] - 1) // 4 + 1
grid_sizes = torch.tensor([latent_frames, 44, 80])
init_kwargs = {
"num_gpus": BASE_PARAMS["num_gpus"],
"use_fsdp_inference": True,
"dit_cpu_offload": False,
"vae_cpu_offload": False,
"text_encoder_cpu_offload": True,
"pin_cpu_memory": True,
}
generation_kwargs = {
"num_inference_steps": num_inference_steps,
"output_path": output_dir,
"image_path": TEST_IMAGE_PATHS[0],
"height": BASE_PARAMS["height"],
"width": BASE_PARAMS["width"],
"num_frames": BASE_PARAMS["num_frames"],
"seed": BASE_PARAMS["seed"],
"mouse_cond": actions["mouse"].unsqueeze(0),
"keyboard_cond": actions["keyboard"].unsqueeze(0),
"grid_sizes": grid_sizes,
"save_video": True,
}
generator = VideoGenerator.from_pretrained(
model_path=BASE_PARAMS["model_path"], **init_kwargs
)
generator.generate_video(prompt, **generation_kwargs)
if isinstance(generator.executor, MultiprocExecutor):
generator.executor.shutdown()
assert os.path.exists(output_dir), (
f"Output video was not generated at {output_dir}"
)
reference_folder = os.path.join(
script_dir, device_reference_folder, model_id, ATTENTION_BACKEND
)
if not os.path.exists(reference_folder):
logger.error("Reference folder missing")
raise FileNotFoundError(
f"Reference video folder does not exist: {reference_folder}"
)
# Find the matching reference video based on the prompt
reference_video_name = None
for filename in os.listdir(reference_folder):
if filename.endswith(".mp4"):
reference_video_name = filename
break
if not reference_video_name:
logger.error(
f"Reference video not found for model: {model_id} with backend: {ATTENTION_BACKEND}"
)
raise FileNotFoundError("Reference video missing")
reference_video_path = os.path.join(reference_folder, reference_video_name)
generated_video_path = os.path.join(output_dir, output_video_name)
logger.info(
f"Computing SSIM between {reference_video_path} and {generated_video_path}"
)
ssim_values = compute_video_ssim_torchvision(
reference_video_path, generated_video_path, use_ms_ssim=True
)
mean_ssim = ssim_values[0]
logger.info(f"SSIM mean value: {mean_ssim}")
logger.info(f"Writing SSIM results to directory: {output_dir}")
success = write_ssim_results(
output_dir,
ssim_values,
reference_video_path,
generated_video_path,
num_inference_steps,
prompt,
)
if not success:
logger.error("Failed to write SSIM results to file")
min_acceptable_ssim = 0.98
assert mean_ssim >= min_acceptable_ssim, (
f"SSIM value {mean_ssim} is below threshold {min_acceptable_ssim} for {model_id} with backend {ATTENTION_BACKEND}"
)

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