Compare commits
26
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ec199c8b2f | ||
|
|
0dc0617ac5 | ||
|
|
087f5b927f | ||
|
|
9d822d464b | ||
|
|
9f86b28ecb | ||
|
|
95b9cde729 | ||
|
|
54b85d8931 | ||
|
|
922b082cb2 | ||
|
|
fb3bfecd18 | ||
|
|
f01fb88c6a | ||
|
|
e8f298c1ab | ||
|
|
d7c7d23375 | ||
|
|
44808ce145 | ||
|
|
0d5124e092 | ||
|
|
7f71994653 | ||
|
|
2bb3349da1 | ||
|
|
8fe1689968 | ||
|
|
e53730f324 | ||
|
|
7a4fe9086a | ||
|
|
d277361aae | ||
|
|
734a54e7a9 | ||
|
|
91364982df | ||
|
|
50145e4fcb | ||
|
|
4112507e99 | ||
|
|
424fc2b4ae | ||
|
|
e6066223e6 |
@@ -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
@@ -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.")
|
||||
@@ -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
@@ -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();
|
||||
|
||||
@@ -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 . .
|
||||
|
||||
|
||||
@@ -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 . .
|
||||
|
||||
|
||||
@@ -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 . .
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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[@]}"
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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
|
||||
@@ -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():
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||
@@ -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())),
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)",
|
||||
|
||||
@@ -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", ""),
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
@@ -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
|
||||
@@ -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}")
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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")
|
||||
}
|
||||
|
||||
@@ -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
|
||||
@@ -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]
|
||||
@@ -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...")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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 = []
|
||||
|
||||
@@ -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__)
|
||||
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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",
|
||||
|
||||
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
BIN
Binary file not shown.
@@ -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
Reference in New Issue
Block a user