[diffusion] perf: replace Conv3d with reshape + F.linear in PatchEmbed (#21014)

This commit is contained in:
AichenF
2026-04-07 09:12:59 +08:00
committed by GitHub
parent ae38b24cc3
commit 5e2b0f860c
3 changed files with 528 additions and 4 deletions
@@ -6,6 +6,7 @@ import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from diffusers.models.embeddings import (
CombinedTimestepGuidanceTextProjEmbeddings as _CombinedTimestepGuidanceTextProjEmbeddings,
)
@@ -58,12 +59,13 @@ class PatchEmbed(nn.Module):
prefix: str = "",
):
super().__init__()
# Convert patch_size to 2-tuple
if isinstance(patch_size, list | tuple):
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:
patch_size = (patch_size, patch_size)
patch_size = (1, patch_size, patch_size)
self.patch_size = patch_size
self.flatten = flatten
@@ -79,9 +81,32 @@ class PatchEmbed(nn.Module):
self.norm = norm_layer(embed_dim) if norm_layer else nn.Identity()
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)
if self.flatten:
x = x.flatten(2).transpose(1, 2) # BCHW -> BNC
x = x.flatten(2).transpose(1, 2)
x = self.norm(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"])