[diffusion] perf: replace Conv3d with reshape + F.linear in PatchEmbed (#21014)
This commit is contained in:
@@ -6,6 +6,7 @@ import math
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
import torch.nn as nn
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
from diffusers.models.embeddings import (
|
from diffusers.models.embeddings import (
|
||||||
CombinedTimestepGuidanceTextProjEmbeddings as _CombinedTimestepGuidanceTextProjEmbeddings,
|
CombinedTimestepGuidanceTextProjEmbeddings as _CombinedTimestepGuidanceTextProjEmbeddings,
|
||||||
)
|
)
|
||||||
@@ -58,12 +59,13 @@ class PatchEmbed(nn.Module):
|
|||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
# Convert patch_size to 2-tuple
|
|
||||||
if isinstance(patch_size, list | tuple):
|
if isinstance(patch_size, list | tuple):
|
||||||
if len(patch_size) == 1:
|
if len(patch_size) == 1:
|
||||||
patch_size = (patch_size[0], patch_size[0])
|
patch_size = (1, patch_size[0], patch_size[0])
|
||||||
|
elif len(patch_size) == 2:
|
||||||
|
patch_size = (1, patch_size[0], patch_size[1])
|
||||||
else:
|
else:
|
||||||
patch_size = (patch_size, patch_size)
|
patch_size = (1, patch_size, patch_size)
|
||||||
|
|
||||||
self.patch_size = patch_size
|
self.patch_size = patch_size
|
||||||
self.flatten = flatten
|
self.flatten = flatten
|
||||||
@@ -79,9 +81,32 @@ class PatchEmbed(nn.Module):
|
|||||||
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
|
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
|
||||||
|
|
||||||
def forward(self, x):
|
def forward(self, x):
|
||||||
|
if x.dim() == 5:
|
||||||
|
B, C, T, H, W = x.shape
|
||||||
|
pt, ph, pw = self.patch_size
|
||||||
|
|
||||||
|
if T % pt == 0 and H % ph == 0 and W % pw == 0:
|
||||||
|
T_ = T // pt
|
||||||
|
H_ = H // ph
|
||||||
|
W_ = W // pw
|
||||||
|
|
||||||
|
x = x.reshape(B, C, T_, pt, H_, ph, W_, pw)
|
||||||
|
x = x.permute(0, 2, 4, 6, 1, 3, 5, 7).contiguous()
|
||||||
|
x = x.reshape(B, T_ * H_ * W_, C * pt * ph * pw)
|
||||||
|
|
||||||
|
w = self.proj.weight.reshape(self.proj.weight.shape[0], -1)
|
||||||
|
x = F.linear(x, w, self.proj.bias) # [B, T'*H'*W', embed_dim]
|
||||||
|
|
||||||
|
if not self.flatten:
|
||||||
|
x = x.reshape(B, T_, H_, W_, -1).permute(0, 4, 1, 2, 3).contiguous()
|
||||||
|
|
||||||
|
x = self.norm(x)
|
||||||
|
return x
|
||||||
|
|
||||||
|
# Fallback to Conv3d for non-5D input or indivisible spatial dims.
|
||||||
x = self.proj(x)
|
x = self.proj(x)
|
||||||
if self.flatten:
|
if self.flatten:
|
||||||
x = x.flatten(2).transpose(1, 2) # BCHW -> BNC
|
x = x.flatten(2).transpose(1, 2)
|
||||||
x = self.norm(x)
|
x = self.norm(x)
|
||||||
return x
|
return x
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,207 @@
|
|||||||
|
"""
|
||||||
|
Benchmark: Conv3d vs reshape + F.linear PatchEmbed.
|
||||||
|
|
||||||
|
Matches the real e2e pipeline conditions:
|
||||||
|
- Conv3d weights are FP32 (no dtype passed to PatchEmbed.__init__)
|
||||||
|
- Input latents are BF16 (cast by denoising loop)
|
||||||
|
- torch.autocast(dtype=bf16) wraps the forward pass
|
||||||
|
- .flatten(2).transpose(1, 2) follows PatchEmbed (wanvideo.py:1008)
|
||||||
|
|
||||||
|
Uses CUDA events for accurate GPU timing. Each case runs warmup iterations
|
||||||
|
followed by timed iterations, reports median latency and speedup.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
python bench_patch_embed.py
|
||||||
|
python bench_patch_embed.py --warmup 20 --iters 100
|
||||||
|
"""
|
||||||
|
|
||||||
|
import argparse
|
||||||
|
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
|
||||||
|
class PatchEmbed3D(nn.Module):
|
||||||
|
"""Conv3d-based PatchEmbed (upstream/main)."""
|
||||||
|
|
||||||
|
def __init__(self, patch_size, in_chans, embed_dim, flatten=True, bias=True):
|
||||||
|
super().__init__()
|
||||||
|
if isinstance(patch_size, list | tuple):
|
||||||
|
if len(patch_size) == 1:
|
||||||
|
patch_size = (patch_size[0], patch_size[0])
|
||||||
|
else:
|
||||||
|
patch_size = (patch_size, patch_size)
|
||||||
|
self.patch_size = patch_size
|
||||||
|
self.flatten = flatten
|
||||||
|
self.proj = nn.Conv3d(
|
||||||
|
in_chans,
|
||||||
|
embed_dim,
|
||||||
|
kernel_size=patch_size,
|
||||||
|
stride=patch_size,
|
||||||
|
bias=bias,
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
x = self.proj(x)
|
||||||
|
if self.flatten:
|
||||||
|
x = x.flatten(2).transpose(1, 2)
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
class PatchEmbed(nn.Module):
|
||||||
|
"""Reshape + F.linear PatchEmbed (opt_krea)."""
|
||||||
|
|
||||||
|
def __init__(self, patch_size, in_chans, embed_dim, flatten=True, bias=True):
|
||||||
|
super().__init__()
|
||||||
|
if isinstance(patch_size, list | tuple):
|
||||||
|
if len(patch_size) == 1:
|
||||||
|
patch_size = (1, patch_size[0], patch_size[0])
|
||||||
|
elif len(patch_size) == 2:
|
||||||
|
patch_size = (1, patch_size[0], patch_size[1])
|
||||||
|
else:
|
||||||
|
patch_size = (1, patch_size, patch_size)
|
||||||
|
self.patch_size = patch_size
|
||||||
|
self.flatten = flatten
|
||||||
|
self.proj = nn.Conv3d(
|
||||||
|
in_chans,
|
||||||
|
embed_dim,
|
||||||
|
kernel_size=patch_size,
|
||||||
|
stride=patch_size,
|
||||||
|
bias=bias,
|
||||||
|
)
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
B, C, T, H, W = x.shape
|
||||||
|
pt, ph, pw = self.patch_size
|
||||||
|
T_ = T // pt
|
||||||
|
H_ = H // ph
|
||||||
|
W_ = W // pw
|
||||||
|
x = x.reshape(B, C, T_, pt, H_, ph, W_, pw)
|
||||||
|
x = x.permute(0, 2, 4, 6, 1, 3, 5, 7).contiguous()
|
||||||
|
x = x.reshape(B, T_ * H_ * W_, C * pt * ph * pw)
|
||||||
|
w = self.proj.weight.reshape(self.proj.weight.shape[0], -1)
|
||||||
|
x = F.linear(x, w, self.proj.bias)
|
||||||
|
if not self.flatten:
|
||||||
|
x = x.reshape(B, T_, H_, W_, -1).permute(0, 4, 1, 2, 3).contiguous()
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
def _copy_weights(src, dst):
|
||||||
|
dst.proj.weight.data.copy_(src.proj.weight.data)
|
||||||
|
if src.proj.bias is not None:
|
||||||
|
dst.proj.bias.data.copy_(src.proj.bias.data)
|
||||||
|
|
||||||
|
|
||||||
|
def bench_one(fn, warmup, iters):
|
||||||
|
"""Returns list of per-iteration latencies in ms using CUDA events."""
|
||||||
|
for _ in range(warmup):
|
||||||
|
fn()
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
|
||||||
|
times = []
|
||||||
|
for _ in range(iters):
|
||||||
|
start = torch.cuda.Event(enable_timing=True)
|
||||||
|
end = torch.cuda.Event(enable_timing=True)
|
||||||
|
start.record()
|
||||||
|
fn()
|
||||||
|
end.record()
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
times.append(start.elapsed_time(end))
|
||||||
|
return times
|
||||||
|
|
||||||
|
|
||||||
|
# Real latent shapes: T = (num_frames-1)//4+1, H = height//8, W = width//8
|
||||||
|
# (name, patch_size, in_chans, embed_dim, flatten, B, T, H, W)
|
||||||
|
BENCH_CASES = [
|
||||||
|
# Wan2.1-I2V-14B: 480x832
|
||||||
|
("Wan-21f-480x832", (1, 2, 2), 16, 5120, False, 1, 6, 60, 104), # 21 frames
|
||||||
|
("Wan-41f-480x832", (1, 2, 2), 16, 5120, False, 1, 11, 60, 104), # 41 frames
|
||||||
|
("Wan-81f-480x832", (1, 2, 2), 16, 5120, False, 1, 21, 60, 104), # 81 frames
|
||||||
|
("Wan-101f-480x832", (1, 2, 2), 16, 5120, False, 1, 26, 60, 104), # 101 frames
|
||||||
|
# Wan2.1-I2V-14B: 720x1280
|
||||||
|
("Wan-21f-720x1280", (1, 2, 2), 16, 5120, False, 1, 6, 90, 160), # 21 frames 720p
|
||||||
|
("Wan-41f-720x1280", (1, 2, 2), 16, 5120, False, 1, 11, 90, 160), # 41 frames 720p
|
||||||
|
# HunyuanVideo
|
||||||
|
("HunYuan-21f-480x832", (1, 2, 2), 16, 3072, True, 1, 6, 60, 104),
|
||||||
|
("HunYuan-41f-480x832", (1, 2, 2), 16, 3072, True, 1, 11, 60, 104),
|
||||||
|
]
|
||||||
|
|
||||||
|
|
||||||
|
def main():
|
||||||
|
parser = argparse.ArgumentParser(
|
||||||
|
description="Benchmark PatchEmbed: Conv3d vs F.linear"
|
||||||
|
)
|
||||||
|
parser.add_argument("--warmup", type=int, default=10)
|
||||||
|
parser.add_argument("--iters", type=int, default=50)
|
||||||
|
args = parser.parse_args()
|
||||||
|
|
||||||
|
device = "cuda"
|
||||||
|
|
||||||
|
# ── Real pipeline conditions ──────────────────────────────────────────
|
||||||
|
# 1. Weights are FP32 (PatchEmbed.__init__ has no dtype arg in real code)
|
||||||
|
# 2. Input is BF16 (latents.to(target_dtype) in denoising loop)
|
||||||
|
# 3. torch.autocast(dtype=bf16) wraps the denoising loop
|
||||||
|
# 4. .flatten(2).transpose(1, 2) follows PatchEmbed (wanvideo.py:1008)
|
||||||
|
# ──────────────────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
header = f"{'Case':<25} {'Conv3d(ms)':>10} {'F.linear(ms)':>12} {'Speedup':>8}"
|
||||||
|
print("Real pipeline conditions: FP32 weights, BF16 input, autocast(bf16)")
|
||||||
|
print(header)
|
||||||
|
print("-" * len(header))
|
||||||
|
|
||||||
|
for name, patch_size, in_chans, embed_dim, flatten, B, T, H, W in BENCH_CASES:
|
||||||
|
torch.manual_seed(42)
|
||||||
|
|
||||||
|
# FP32 weights – matches real model init (no dtype passed)
|
||||||
|
conv_model = (
|
||||||
|
PatchEmbed3D(
|
||||||
|
patch_size,
|
||||||
|
in_chans,
|
||||||
|
embed_dim,
|
||||||
|
flatten,
|
||||||
|
)
|
||||||
|
.to(device)
|
||||||
|
.eval()
|
||||||
|
)
|
||||||
|
lin_model = (
|
||||||
|
PatchEmbed(
|
||||||
|
patch_size,
|
||||||
|
in_chans,
|
||||||
|
embed_dim,
|
||||||
|
flatten,
|
||||||
|
)
|
||||||
|
.to(device)
|
||||||
|
.eval()
|
||||||
|
)
|
||||||
|
_copy_weights(conv_model, lin_model)
|
||||||
|
|
||||||
|
# BF16 input – matches real latent dtype
|
||||||
|
x = torch.randn(B, in_chans, T, H, W, device=device, dtype=torch.bfloat16)
|
||||||
|
|
||||||
|
# Include the .flatten(2).transpose(1, 2) that follows PatchEmbed
|
||||||
|
# in WanTransformer3DModel.forward (wanvideo.py:1008)
|
||||||
|
def conv_fn():
|
||||||
|
out = conv_model(x)
|
||||||
|
return out.flatten(2).transpose(1, 2)
|
||||||
|
|
||||||
|
def lin_fn():
|
||||||
|
out = lin_model(x)
|
||||||
|
return out.flatten(2).transpose(1, 2)
|
||||||
|
|
||||||
|
# autocast(bf16) – matches real denoising loop (denoising.py:1016)
|
||||||
|
with torch.no_grad(), torch.autocast(device_type="cuda", dtype=torch.bfloat16):
|
||||||
|
t_conv = bench_one(conv_fn, args.warmup, args.iters)
|
||||||
|
t_lin = bench_one(lin_fn, args.warmup, args.iters)
|
||||||
|
|
||||||
|
med_conv = sorted(t_conv)[len(t_conv) // 2]
|
||||||
|
med_lin = sorted(t_lin)[len(t_lin) // 2]
|
||||||
|
speedup = med_conv / med_lin if med_lin > 0 else float("inf")
|
||||||
|
|
||||||
|
print(f"{name:<25} {med_conv:>10.3f} {med_lin:>12.3f} {speedup:>7.2f}x")
|
||||||
|
|
||||||
|
print()
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -0,0 +1,292 @@
|
|||||||
|
"""
|
||||||
|
Test that the optimized PatchEmbed (reshape + F.linear) is equivalent
|
||||||
|
to the original Conv3d-based PatchEmbed from upstream/main.
|
||||||
|
|
||||||
|
The opt_krea branch replaces Conv3d forward with manual
|
||||||
|
reshape + permute + F.linear for 5D input. This is valid because
|
||||||
|
Conv3d with stride==kernel_size is a non-overlapping patch extraction
|
||||||
|
followed by linear projection, which is exactly what the manual path does.
|
||||||
|
|
||||||
|
We disable TF32 so cuDNN (Conv3d) and cuBLAS (F.linear) both use
|
||||||
|
full FP32 precision, enabling strict numerical comparison.
|
||||||
|
"""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
import torch.nn as nn
|
||||||
|
import torch.nn.functional as F
|
||||||
|
|
||||||
|
|
||||||
|
class PatchEmbed3D(nn.Module):
|
||||||
|
"""PatchEmbed from upstream/main: uses Conv3d directly."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self, patch_size, in_chans, embed_dim, flatten=True, bias=True, dtype=None
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
if isinstance(patch_size, list | tuple):
|
||||||
|
if len(patch_size) == 1:
|
||||||
|
patch_size = (patch_size[0], patch_size[0])
|
||||||
|
else:
|
||||||
|
patch_size = (patch_size, patch_size)
|
||||||
|
self.patch_size = patch_size
|
||||||
|
self.flatten = flatten
|
||||||
|
self.proj = nn.Conv3d(
|
||||||
|
in_chans,
|
||||||
|
embed_dim,
|
||||||
|
kernel_size=patch_size,
|
||||||
|
stride=patch_size,
|
||||||
|
bias=bias,
|
||||||
|
dtype=dtype,
|
||||||
|
)
|
||||||
|
self.norm = nn.Identity()
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
x = self.proj(x)
|
||||||
|
if self.flatten:
|
||||||
|
x = x.flatten(2).transpose(1, 2)
|
||||||
|
x = self.norm(x)
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
class PatchEmbed(nn.Module):
|
||||||
|
"""PatchEmbed from opt_krea: replaces Conv3d with reshape + F.linear for 5D input."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self, patch_size, in_chans, embed_dim, flatten=True, bias=True, dtype=None
|
||||||
|
):
|
||||||
|
super().__init__()
|
||||||
|
if isinstance(patch_size, list | tuple):
|
||||||
|
if len(patch_size) == 1:
|
||||||
|
patch_size = (1, patch_size[0], patch_size[0])
|
||||||
|
elif len(patch_size) == 2:
|
||||||
|
patch_size = (1, patch_size[0], patch_size[1])
|
||||||
|
else:
|
||||||
|
patch_size = (1, patch_size, patch_size)
|
||||||
|
self.patch_size = patch_size
|
||||||
|
self.flatten = flatten
|
||||||
|
self.proj = nn.Conv3d(
|
||||||
|
in_chans,
|
||||||
|
embed_dim,
|
||||||
|
kernel_size=patch_size,
|
||||||
|
stride=patch_size,
|
||||||
|
bias=bias,
|
||||||
|
dtype=dtype,
|
||||||
|
)
|
||||||
|
self.norm = nn.Identity()
|
||||||
|
|
||||||
|
def forward(self, x):
|
||||||
|
if x.dim() == 5:
|
||||||
|
B, C, T, H, W = x.shape
|
||||||
|
pt, ph, pw = self.patch_size
|
||||||
|
T_ = T // pt
|
||||||
|
H_ = H // ph
|
||||||
|
W_ = W // pw
|
||||||
|
x = x.reshape(B, C, T_, pt, H_, ph, W_, pw)
|
||||||
|
x = x.permute(0, 2, 4, 6, 1, 3, 5, 7).contiguous()
|
||||||
|
x = x.reshape(B, T_ * H_ * W_, C * pt * ph * pw)
|
||||||
|
w = self.proj.weight.reshape(self.proj.weight.shape[0], -1)
|
||||||
|
x = F.linear(x, w, self.proj.bias)
|
||||||
|
if not self.flatten:
|
||||||
|
x = x.reshape(B, T_, H_, W_, -1).permute(0, 4, 1, 2, 3).contiguous()
|
||||||
|
x = self.norm(x)
|
||||||
|
return x
|
||||||
|
x = self.proj(x)
|
||||||
|
if self.flatten:
|
||||||
|
x = x.flatten(2).transpose(1, 2)
|
||||||
|
x = self.norm(x)
|
||||||
|
return x
|
||||||
|
|
||||||
|
|
||||||
|
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
|
||||||
|
|
||||||
|
|
||||||
|
def _copy_weights(src, dst):
|
||||||
|
dst.proj.weight.data.copy_(src.proj.weight.data)
|
||||||
|
if src.proj.bias is not None:
|
||||||
|
dst.proj.bias.data.copy_(src.proj.bias.data)
|
||||||
|
|
||||||
|
|
||||||
|
def _run_equivalence(
|
||||||
|
patch_size,
|
||||||
|
in_chans,
|
||||||
|
embed_dim,
|
||||||
|
flatten,
|
||||||
|
bias,
|
||||||
|
weight_dtype,
|
||||||
|
input_dtype,
|
||||||
|
B,
|
||||||
|
T,
|
||||||
|
H,
|
||||||
|
W,
|
||||||
|
atol,
|
||||||
|
rtol,
|
||||||
|
):
|
||||||
|
"""Helper: build both models with shared weights, run forward, compare.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
weight_dtype: dtype for Conv3d weights (None = FP32).
|
||||||
|
input_dtype: dtype for the input tensor (None = FP32).
|
||||||
|
"""
|
||||||
|
torch.manual_seed(42)
|
||||||
|
main = (
|
||||||
|
PatchEmbed3D(patch_size, in_chans, embed_dim, flatten, bias, dtype=weight_dtype)
|
||||||
|
.to(DEVICE)
|
||||||
|
.eval()
|
||||||
|
)
|
||||||
|
opt = (
|
||||||
|
PatchEmbed(patch_size, in_chans, embed_dim, flatten, bias, dtype=weight_dtype)
|
||||||
|
.to(DEVICE)
|
||||||
|
.eval()
|
||||||
|
)
|
||||||
|
_copy_weights(main, opt)
|
||||||
|
|
||||||
|
x = torch.randn(
|
||||||
|
B, in_chans, T, H, W, device=DEVICE, dtype=input_dtype or torch.float32
|
||||||
|
)
|
||||||
|
with torch.no_grad():
|
||||||
|
out_main = main(x)
|
||||||
|
out_opt = opt(x)
|
||||||
|
|
||||||
|
assert (
|
||||||
|
out_main.shape == out_opt.shape
|
||||||
|
), f"Shape mismatch: {out_main.shape} vs {out_opt.shape}"
|
||||||
|
assert (
|
||||||
|
out_main.dtype == out_opt.dtype
|
||||||
|
), f"Dtype mismatch: {out_main.dtype} vs {out_opt.dtype}"
|
||||||
|
torch.testing.assert_close(out_main, out_opt, atol=atol, rtol=rtol)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture(autouse=True)
|
||||||
|
def _disable_tf32():
|
||||||
|
prev_cudnn = torch.backends.cudnn.allow_tf32
|
||||||
|
prev_matmul = torch.backends.cuda.matmul.allow_tf32
|
||||||
|
torch.backends.cudnn.allow_tf32 = False
|
||||||
|
torch.backends.cuda.matmul.allow_tf32 = False
|
||||||
|
yield
|
||||||
|
torch.backends.cudnn.allow_tf32 = prev_cudnn
|
||||||
|
torch.backends.cuda.matmul.allow_tf32 = prev_matmul
|
||||||
|
|
||||||
|
|
||||||
|
# ── Wan2.1 / Wan2.2 / CausalWan / Helios ────────────────────────────────────
|
||||||
|
# patch_size=(1,2,2), in_channels=16, embed_dim=5120, flatten=False
|
||||||
|
# Real usage: weight=FP32 (no dtype passed), input=BF16 from VAE latent
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"dtype,atol,rtol",
|
||||||
|
[
|
||||||
|
(None, 1e-4, 1e-4), # weight=FP32, input=FP32
|
||||||
|
(torch.bfloat16, 1e-2, 1e-2), # weight=BF16, input=BF16
|
||||||
|
(torch.float16, 1e-2, 1e-2), # weight=FP16, input=FP16
|
||||||
|
],
|
||||||
|
ids=["fp32", "bf16", "fp16"],
|
||||||
|
)
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"B,T,H,W",
|
||||||
|
[
|
||||||
|
(1, 21, 60, 104), # 480p typical
|
||||||
|
(2, 9, 40, 64), # smaller resolution, batch=2
|
||||||
|
(1, 33, 90, 160), # 720p longer video
|
||||||
|
],
|
||||||
|
ids=["480p-B1", "small-B2", "720p-B1"],
|
||||||
|
)
|
||||||
|
def test_wan_helios(dtype, atol, rtol, B, T, H, W):
|
||||||
|
_run_equivalence(
|
||||||
|
patch_size=(1, 2, 2),
|
||||||
|
in_chans=16,
|
||||||
|
embed_dim=5120,
|
||||||
|
flatten=False,
|
||||||
|
bias=True,
|
||||||
|
weight_dtype=dtype,
|
||||||
|
input_dtype=dtype,
|
||||||
|
B=B,
|
||||||
|
T=T,
|
||||||
|
H=H,
|
||||||
|
W=W,
|
||||||
|
atol=atol,
|
||||||
|
rtol=rtol,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ── HunyuanVideo ─────────────────────────────────────────────────────────────
|
||||||
|
# patch_size=[1,2,2] (list!), in_channels=16, embed_dim=3072, flatten=True
|
||||||
|
# Real usage: dtype passed to PatchEmbed, so weight & input share same dtype
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"dtype,atol,rtol",
|
||||||
|
[
|
||||||
|
(None, 1e-4, 1e-4), # weight=FP32, input=FP32
|
||||||
|
(torch.bfloat16, 1e-2, 1e-2), # weight=BF16, input=BF16
|
||||||
|
(torch.float16, 1e-2, 1e-2), # weight=FP16, input=FP16
|
||||||
|
],
|
||||||
|
ids=["fp32", "bf16", "fp16"],
|
||||||
|
)
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
"B,T,H,W",
|
||||||
|
[
|
||||||
|
(1, 21, 60, 104),
|
||||||
|
(2, 9, 40, 64),
|
||||||
|
],
|
||||||
|
ids=["480p-B1", "small-B2"],
|
||||||
|
)
|
||||||
|
def test_hunyuanvideo(dtype, atol, rtol, B, T, H, W):
|
||||||
|
_run_equivalence(
|
||||||
|
patch_size=[1, 2, 2],
|
||||||
|
in_chans=16,
|
||||||
|
embed_dim=3072,
|
||||||
|
flatten=True,
|
||||||
|
bias=True,
|
||||||
|
weight_dtype=dtype,
|
||||||
|
input_dtype=dtype,
|
||||||
|
B=B,
|
||||||
|
T=T,
|
||||||
|
H=H,
|
||||||
|
W=W,
|
||||||
|
atol=atol,
|
||||||
|
rtol=rtol,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# ── No-bias variants ─────────────────────────────────────────────────────────
|
||||||
|
|
||||||
|
|
||||||
|
def test_wan_no_bias():
|
||||||
|
_run_equivalence(
|
||||||
|
patch_size=(1, 2, 2),
|
||||||
|
in_chans=16,
|
||||||
|
embed_dim=5120,
|
||||||
|
flatten=False,
|
||||||
|
bias=False,
|
||||||
|
weight_dtype=None,
|
||||||
|
input_dtype=None,
|
||||||
|
B=1,
|
||||||
|
T=21,
|
||||||
|
H=60,
|
||||||
|
W=104,
|
||||||
|
atol=1e-4,
|
||||||
|
rtol=1e-4,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_hunyuanvideo_no_bias():
|
||||||
|
_run_equivalence(
|
||||||
|
patch_size=[1, 2, 2],
|
||||||
|
in_chans=16,
|
||||||
|
embed_dim=3072,
|
||||||
|
flatten=True,
|
||||||
|
bias=False,
|
||||||
|
weight_dtype=None,
|
||||||
|
input_dtype=None,
|
||||||
|
B=1,
|
||||||
|
T=21,
|
||||||
|
H=60,
|
||||||
|
W=104,
|
||||||
|
atol=1e-4,
|
||||||
|
rtol=1e-4,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
pytest.main([__file__, "-v", "--tb=short"])
|
||||||
Reference in New Issue
Block a user