Support 5D video VAE latents (Wan) in patchify/unpatchify

Video VAEs output [B, C, T, H, W]. Patchify now merges T into batch,
processes all frames, and unpatchify restores the 5D shape.
Updated all callers to pass extra_shape through.
This commit is contained in:
facok
2026-03-21 16:19:03 +08:00
parent a336d4ca95
commit 0d82606fb0
6 changed files with 37 additions and 26 deletions
+3 -3
View File
@@ -90,7 +90,7 @@ def calibrate(vae, num_colors=512, image_size=512, batch_size=8):
latent = vae.encode(imgs[:, :, :, :3])
# Patchify → [B', L, D]
patches, _, _ = patchify(latent)
patches, _, _, _ = patchify(latent)
# Average across patches → [B', D]
avg = patches.mean(dim=1).cpu()
@@ -104,7 +104,7 @@ def calibrate(vae, num_colors=512, image_size=512, batch_size=8):
for k in range(1, actual_batch):
single = imgs[k:k+1, :, :, :3]
lat = vae.encode(single)
p, _, _ = patchify(lat)
p, _, _, _ = patchify(lat)
vectors.append(p.mean(dim=1).cpu().squeeze(0))
pbar.update(actual_batch)
@@ -137,7 +137,7 @@ def calibrate(vae, num_colors=512, image_size=512, batch_size=8):
img[0, :, :, 1] = g
img[0, :, :, 2] = b
latent = vae.encode(img[:, :, :, :3])
patches, _, _ = patchify(latent)
patches, _, _, _ = patchify(latent)
avg = patches.mean(dim=1).cpu().squeeze(0) # [64]
# Project to LCS
lcs_coord = (avg - mean) @ basis # [3]
+25 -14
View File
@@ -4,31 +4,42 @@ from einops import rearrange
def patchify(x):
"""Convert latent [B, C, H, W] or [B, C, ..., H, W] → patch sequence [B, L, C*4].
"""Convert latent [B, C, H, W] or [B, C, T, H, W] → patch sequence [B, L, C*4].
For video VAEs (5D input with T frames), merges T into the batch dimension
so all frames are patchified together. The original shape is returned as
extra_shape for unpatchify to restore.
For video VAEs (5D+ input), collapses all dimensions between C and H×W
by reshaping to 4D. Only the last frame/slice is used for spatial patchification.
L = (H/2) * (W/2), d = C * 2 * 2.
Returns (patches, h_len, w_len) where h_len=H/2, w_len=W/2.
Returns (patches, h_len, w_len, extra_shape) where extra_shape is None for 4D
or (B_orig, C, T) for 5D.
"""
if x.ndim > 4:
# Video VAE: [B, C, T, H, W] or similar — take first frame
while x.ndim > 4:
x = x[:, :, 0]
extra_shape = None
if x.ndim == 5:
B_orig, C, T, H, W = x.shape
extra_shape = (B_orig, C, T)
# Merge B and T: [B*T, C, H, W]
x = x.permute(0, 2, 1, 3, 4).reshape(B_orig * T, C, H, W)
B, C, H, W = x.shape
h_len = H // 2
w_len = W // 2
patches = rearrange(x, "b c (h ph) (w pw) -> b (h w) (c ph pw)", ph=2, pw=2)
return patches, h_len, w_len
return patches, h_len, w_len, extra_shape
def unpatchify(patches, h_len, w_len):
"""Convert patch sequence [B, L, C*4] → latent [B, C, H, W].
def unpatchify(patches, h_len, w_len, extra_shape=None):
"""Convert patch sequence [B, L, C*4] → latent [B, C, H, W] or [B, C, T, H, W].
Auto-detects channel count from patch dimension: C = D / 4.
h_len, w_len from patchify output.
If extra_shape is provided, restores the 5D video format.
"""
D = patches.shape[-1]
C = D // 4 # patch_size=2×2=4
return rearrange(patches, "b (h w) (c ph pw) -> b c (h ph) (w pw)",
h=h_len, w=w_len, c=C, ph=2, pw=2)
x = rearrange(patches, "b (h w) (c ph pw) -> b c (h ph) (w pw)",
h=h_len, w=w_len, c=C, ph=2, pw=2)
if extra_shape is not None:
B_orig, C_orig, T = extra_shape
# Unmerge B and T: [B_orig*T, C, H, W] → [B_orig, C, T, H, W]
H, W = x.shape[2], x.shape[3]
x = x.reshape(B_orig, T, C, H, W).permute(0, 2, 1, 3, 4)
return x
+2 -2
View File
@@ -162,7 +162,7 @@ def calibrate_sharpness(vae, num_samples: int = 64, image_size: int = 512,
# VAE encode — try batch first, fall back to per-image for video VAEs
latent = vae.encode(imgs_bhwc)
patches, _, _ = patchify(latent)
patches, _, _, _ = patchify(latent)
avg = patches.mean(dim=1).cpu()
if avg.shape[0] == actual_batch:
@@ -176,7 +176,7 @@ def calibrate_sharpness(vae, num_samples: int = 64, image_size: int = 512,
for k in range(1, actual_batch):
single = imgs_bhwc[k:k+1]
lat = vae.encode(single)
p, _, _ = patchify(lat)
p, _, _, _ = patchify(lat)
vectors.append(p.mean(dim=1).cpu().squeeze(0))
blur_labels.append(blur_sigma)
+4 -4
View File
@@ -49,7 +49,7 @@ def _build_post_cfg_fn(lcs_data, target_colors_hsl, strength, mode, start_step,
raw = denoised / SCALE_FACTOR + SHIFT_FACTOR # [B, 16, H, W]
# Patchify
patches, h_len, w_len = patchify(raw) # [B, L, 64]
patches, h_len, w_len, extra_shape = patchify(raw)
# Project to LCS
projection = (patches - mu) @ B_mat # [B, L, 3]
@@ -155,7 +155,7 @@ def _build_post_cfg_fn(lcs_data, target_colors_hsl, strength, mode, start_step,
patches_new = new_projection @ B_mat.T + mu + residual # [B, L, 64]
# Unpatchify
raw_new = unpatchify(patches_new, h_len, w_len) # [B, 16, H, W]
raw_new = unpatchify(patches_new, h_len, w_len, extra_shape)
# Convert back to process_in space
modified = (raw_new - SHIFT_FACTOR) * SCALE_FACTOR
@@ -330,7 +330,7 @@ def _build_tone_fn(lcs_data, contrast, brightness, saturation, color_temperature
raw = denoised / SCALE_FACTOR + SHIFT_FACTOR # [B, 16, H, W]
# Patchify
patches, h_len, w_len = patchify(raw) # [B, L, 64]
patches, h_len, w_len, extra_shape = patchify(raw)
# Project to LCS
projection = (patches - mu) @ B_mat # [B, L, 3]
@@ -405,7 +405,7 @@ def _build_tone_fn(lcs_data, contrast, brightness, saturation, color_temperature
patches_new = new_projection @ B_mat.T + mu + residual # [B, L, 64]
# Unpatchify
raw_new = unpatchify(patches_new, h_len, w_len) # [B, 16, H, W]
raw_new = unpatchify(patches_new, h_len, w_len, extra_shape)
# Convert back to process_in space
modified = (raw_new - SHIFT_FACTOR) * SCALE_FACTOR
+1 -1
View File
@@ -31,7 +31,7 @@ def _latent_to_color_preview(samples, lcs_data, sigma, upscale=8):
ld = lcs_data.to(device, dtype)
raw = samples / SCALE_FACTOR + SHIFT_FACTOR
patches, h_len, w_len = patchify(raw)
patches, h_len, w_len, _ = patchify(raw)
projection = (patches - ld.mean) @ ld.basis
alpha_t, beta_t = get_alpha_beta(sigma, device=device)
+2 -2
View File
@@ -135,7 +135,7 @@ def _build_sharpness_fn(sharpness_data, strength, start_step, end_step, mask):
raw = denoised / SCALE_FACTOR + SHIFT_FACTOR # [B, 16, H, W]
# Patchify
patches, h_len, w_len = patchify(raw) # [B, L, 64]
patches, h_len, w_len, extra_shape = patchify(raw)
# Apply sharpness edit
if mask is not None:
@@ -147,7 +147,7 @@ def _build_sharpness_fn(sharpness_data, strength, start_step, end_step, mask):
patches_new = patches + ev
# Unpatchify
raw_new = unpatchify(patches_new, h_len, w_len) # [B, 16, H, W]
raw_new = unpatchify(patches_new, h_len, w_len, extra_shape)
# Convert back to process_in space
return ((raw_new - SHIFT_FACTOR) * SCALE_FACTOR).to(dtype)