From 0d82606fb09ff4ef60ed8b4c600f320091be91f0 Mon Sep 17 00:00:00 2001 From: facok <128763816+facok@users.noreply.github.com> Date: Sat, 21 Mar 2026 16:19:03 +0800 Subject: [PATCH] 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. --- core/calibration.py | 6 +++--- core/patchify.py | 39 +++++++++++++++++++++++++-------------- core/sharpness.py | 4 ++-- nodes/intervene.py | 8 ++++---- nodes/observe.py | 2 +- nodes/sharpen.py | 4 ++-- 6 files changed, 37 insertions(+), 26 deletions(-) diff --git a/core/calibration.py b/core/calibration.py index 9958557..adf7014 100644 --- a/core/calibration.py +++ b/core/calibration.py @@ -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] diff --git a/core/patchify.py b/core/patchify.py index 740bf8b..da79e7f 100644 --- a/core/patchify.py +++ b/core/patchify.py @@ -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 diff --git a/core/sharpness.py b/core/sharpness.py index 9d957f9..17c15d5 100644 --- a/core/sharpness.py +++ b/core/sharpness.py @@ -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) diff --git a/nodes/intervene.py b/nodes/intervene.py index d9f6f18..4ebb5e2 100644 --- a/nodes/intervene.py +++ b/nodes/intervene.py @@ -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 diff --git a/nodes/observe.py b/nodes/observe.py index f44a700..ec95108 100644 --- a/nodes/observe.py +++ b/nodes/observe.py @@ -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) diff --git a/nodes/sharpen.py b/nodes/sharpen.py index a89a3c3..37aeab3 100644 --- a/nodes/sharpen.py +++ b/nodes/sharpen.py @@ -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)