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:
+3
-3
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user